use super::blob::{Blob, BlobDecode, BlobReader, BlobType, DecompressPool};
use super::block::{HeaderBlock, PrimitiveBlock};
use super::elements::Element;
use super::file_reader::FileReader;
use super::pipeline::PipelineConfig;
use crate::blob_meta::BlobFilter;
use crate::error::{ErrorKind, Result, new_error};
use std::collections::VecDeque;
use std::io::Read;
use std::path::Path;
use std::sync::mpsc::{Receiver, sync_channel};
use std::sync::{Arc, Condvar, Mutex, PoisonError};
use std::thread::JoinHandle;
const BLOCK_QUEUE: usize = 8;
#[derive(Clone, Debug)]
pub struct ElementReader<R: Read + Send> {
blob_iter: BlobReader<R>,
header: HeaderBlock,
decode_threads: Option<usize>,
pipeline_config: PipelineConfig,
blob_filter: Option<BlobFilter>,
}
impl<R: Read + Send> ElementReader<R> {
pub fn new(reader: R) -> Result<ElementReader<R>> {
let mut blob_iter = BlobReader::new(reader);
let header = read_header_blob(&mut blob_iter)?;
Ok(ElementReader {
blob_iter,
header,
decode_threads: None,
pipeline_config: PipelineConfig::default(),
blob_filter: None,
})
}
pub fn with_blob_filter(mut self, filter: BlobFilter) -> Self {
self.blob_filter = Some(filter);
self
}
pub fn decode_threads(mut self, n: usize) -> Self {
self.decode_threads = Some(n.max(1));
self
}
pub fn read_ahead(mut self, n: usize) -> Self {
self.pipeline_config.read_ahead = n.max(1);
self
}
pub fn decode_ahead(mut self, n: usize) -> Self {
self.pipeline_config.decode_ahead = n.max(1);
self
}
pub fn header(&self) -> &HeaderBlock {
&self.header
}
#[hotpath::measure]
pub fn for_each<F>(self, mut f: F) -> Result<()>
where
F: for<'a> FnMut(Element<'a>),
{
let Self {
blob_iter, header, ..
} = self;
let is_sorted = header.is_sorted();
let mut last_node_id: i64 = i64::MIN;
let mut buf: Vec<u8> = Vec::new();
let mut st_scratch: Vec<(u32, u32)> = Vec::new();
let mut gr_scratch: Vec<(u32, u32)> = Vec::new();
for blob in blob_iter {
let blob = blob?;
if blob.get_type() != BlobType::OsmData {
continue;
}
blob.decompress_into(&mut buf)?;
let block = PrimitiveBlock::from_vec_with_scratch(
std::mem::take(&mut buf),
&mut st_scratch,
&mut gr_scratch,
)?;
block.for_each_element(|element| {
if is_sorted && let Some(id) = node_id(&element) {
debug_assert!(
id > last_node_id,
"Sort.Type_then_ID violated: node {id} <= previous {last_node_id}"
);
last_node_id = id;
}
f(element);
});
}
Ok(())
}
#[hotpath::measure]
pub fn for_each_pipelined<F>(self, mut f: F) -> Result<()>
where
F: for<'a> FnMut(Element<'a>),
{
let is_sorted = self.header.is_sorted();
let is_history = self.header.has_historical_information();
let mut last_node_id: i64 = i64::MIN;
self.for_each_block_pipelined(|block| {
block.for_each_element(|element| {
if is_sorted && let Some(id) = node_id(&element) {
debug_assert!(
if is_history {
id >= last_node_id
} else {
id > last_node_id
},
"Sort.Type_then_ID violated: node {id} < previous {last_node_id}"
);
last_node_id = id;
}
f(element);
});
Ok(())
})
}
pub fn for_each_block_pipelined<F>(self, f: F) -> Result<()>
where
F: FnMut(PrimitiveBlock) -> Result<()>,
{
super::pipeline::run_pipeline(
self.blob_iter,
self.decode_threads,
self.pipeline_config,
self.blob_filter,
f,
)
}
pub(crate) fn for_each_fused_block<T, X, F>(self, transform: X, consume: F) -> Result<()>
where
T: Send,
X: Fn(PrimitiveBlock) -> std::result::Result<T, String> + Sync,
F: FnMut(T) -> Result<()>,
{
super::pipeline::run_pipeline_fused(
self.blob_iter,
self.decode_threads,
self.pipeline_config,
self.blob_filter,
&transform,
consume,
)
}
pub fn into_blocks_pipelined(self) -> PipelinedBlocks
where
R: 'static,
{
let (tx, rx) = sync_channel(BLOCK_QUEUE);
let blob_iter = self.blob_iter;
let decode_threads = self.decode_threads;
let pipeline_config = self.pipeline_config;
let blob_filter = self.blob_filter;
let handle = std::thread::spawn(move || {
let deliver = |block: PrimitiveBlock| {
tx.send(Ok(block)).map_err(|_| {
new_error(ErrorKind::Io(std::io::Error::other(
"pipeline consumer dropped",
)))
})
};
let result = super::pipeline::run_pipeline(
blob_iter,
decode_threads,
pipeline_config,
blob_filter,
deliver,
);
if let Err(e) = result {
drop(tx.send(Err(e)));
}
});
PipelinedBlocks {
rx: Some(rx),
handle: Some(handle),
}
}
pub fn par_map_reduce<MP, RD, ID, T>(
mut self,
map_op: MP,
identity: ID,
reduce_op: RD,
) -> Result<T>
where
MP: for<'a> Fn(Element<'a>) -> T + Sync + Send,
RD: Fn(T, T) -> T + Sync + Send,
ID: Fn() -> T + Sync + Send,
T: Send,
{
self.blob_iter.set_parse_indexdata(false);
let worker_count = self.decode_threads.unwrap_or_else(|| {
std::thread::available_parallelism()
.map(|n| n.get().saturating_sub(1).max(1))
.unwrap_or(3)
});
par_fold_blobs(
self.blob_iter,
worker_count,
PAR_INFLIGHT_BUDGET,
PAR_BATCH_MAX_BLOBS,
PAR_BATCH_MAX_BYTES,
map_op,
identity,
reduce_op,
)
}
}
#[allow(clippy::too_many_arguments)]
fn par_fold_blobs<R, MP, RD, ID, T>(
blob_iter: BlobReader<R>,
worker_count: usize,
budget_cap: u64,
batch_max_blobs: usize,
batch_max_bytes: u64,
map_op: MP,
identity: ID,
reduce_op: RD,
) -> Result<T>
where
R: Read + Send,
MP: for<'a> Fn(Element<'a>) -> T + Sync + Send,
RD: Fn(T, T) -> T + Sync + Send,
ID: Fn() -> T + Sync + Send,
T: Send,
{
let queue = Arc::new(BatchQueue::new());
let budget = Arc::new(ByteBudget::new(budget_cap));
let pool = DecompressPool::new();
let map_op = &map_op;
let identity = &identity;
let reduce_op = &reduce_op;
std::thread::scope(|scope| -> Result<T> {
let mut cancel = CancelGuard::new(&queue, &budget);
let mut handles = Vec::with_capacity(worker_count);
for _ in 0..worker_count {
let queue = Arc::clone(&queue);
let budget = Arc::clone(&budget);
let pool = Arc::clone(&pool);
handles.push(scope.spawn(move || -> Result<T> {
run_par_worker(&queue, &budget, &pool, map_op, identity, reduce_op)
}));
}
let pump_result = pump_blobs(blob_iter, &queue, &budget, batch_max_blobs, batch_max_bytes);
if pump_result.is_err() {
budget.shutdown();
}
queue.close();
cancel.disarm();
let mut partials: Vec<T> = Vec::with_capacity(worker_count);
let mut worker_err: Option<crate::error::Error> = None;
for handle in handles {
match handle.join() {
Ok(Ok(partial)) => partials.push(partial),
Ok(Err(e)) => {
if worker_err.is_none() {
worker_err = Some(e);
}
}
Err(panic) => std::panic::resume_unwind(panic),
}
}
pump_result?;
if let Some(e) = worker_err {
return Err(e);
}
let mut acc = identity();
for partial in partials {
acc = reduce_op(acc, partial);
}
Ok(acc)
})
}
const PAR_INFLIGHT_BUDGET: u64 = 256 * 1024 * 1024;
const PAR_BATCH_MAX_BLOBS: usize = 64;
const PAR_BATCH_MAX_BYTES: u64 = 4 * 1024 * 1024;
struct Batch {
blobs: Vec<Blob>,
bytes: u64,
}
impl Batch {
fn new() -> Self {
Self {
blobs: Vec::new(),
bytes: 0,
}
}
}
struct BatchCharge<'a> {
budget: &'a ByteBudget,
batch: Batch,
}
impl<'a> BatchCharge<'a> {
fn new(budget: &'a ByteBudget, batch: Batch) -> Self {
Self { budget, batch }
}
fn blobs(&self) -> &[Blob] {
&self.batch.blobs
}
}
impl Drop for BatchCharge<'_> {
fn drop(&mut self) {
let bytes = self.batch.bytes;
self.batch.blobs = Vec::new();
self.budget.release(bytes);
}
}
struct CancelGuard<'a> {
queue: &'a BatchQueue,
budget: &'a ByteBudget,
armed: bool,
}
impl<'a> CancelGuard<'a> {
fn new(queue: &'a BatchQueue, budget: &'a ByteBudget) -> Self {
Self {
queue,
budget,
armed: true,
}
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for CancelGuard<'_> {
fn drop(&mut self) {
if self.armed {
self.budget.shutdown();
self.queue.close();
}
}
}
fn pump_blobs<R: Read + Send>(
blob_iter: BlobReader<R>,
queue: &BatchQueue,
budget: &ByteBudget,
batch_max_blobs: usize,
batch_max_bytes: u64,
) -> Result<()> {
let mut batch = Batch::new();
for blob_result in blob_iter {
let blob = blob_result?;
if budget.is_shutdown() {
break;
}
if blob.get_type() != BlobType::OsmData {
continue; }
let weight = blob.retained_len();
if !budget.acquire(weight) {
break; }
if !batch.blobs.is_empty() && batch.bytes.saturating_add(weight) > batch_max_bytes {
queue.push(std::mem::replace(&mut batch, Batch::new()));
}
batch.bytes += weight;
batch.blobs.push(blob);
if batch.blobs.len() >= batch_max_blobs {
queue.push(std::mem::replace(&mut batch, Batch::new()));
}
}
if !batch.blobs.is_empty() {
queue.push(batch);
}
Ok(())
}
fn run_par_worker<MP, RD, ID, T>(
queue: &BatchQueue,
budget: &ByteBudget,
pool: &Arc<DecompressPool>,
map_op: &MP,
identity: &ID,
reduce_op: &RD,
) -> Result<T>
where
MP: for<'a> Fn(Element<'a>) -> T,
RD: Fn(T, T) -> T,
ID: Fn() -> T,
{
let mut cancel = CancelGuard::new(queue, budget);
let mut st_scratch: Vec<(u32, u32)> = Vec::new();
let mut gr_scratch: Vec<(u32, u32)> = Vec::new();
let mut acc = identity();
while let Some(batch) = queue.pop() {
if budget.is_shutdown() {
drop(BatchCharge::new(budget, batch));
break;
}
let charge = BatchCharge::new(budget, batch);
let mut decode_err: Option<crate::error::Error> = None;
for blob in charge.blobs() {
match blob.to_primitiveblock_inline_with_scratch(pool, &mut st_scratch, &mut gr_scratch)
{
Ok(block) => {
for element in block.elements() {
acc = reduce_op(acc, map_op(element));
}
}
Err(e) => {
decode_err = Some(e);
break;
}
}
}
if let Some(e) = decode_err {
budget.shutdown();
queue.close();
cancel.disarm(); drop(charge); return Err(e);
}
}
cancel.disarm();
Ok(acc)
}
struct BatchQueue {
inner: Mutex<BatchQueueState>,
cond: Condvar,
}
struct BatchQueueState {
batches: VecDeque<Batch>,
closed: bool,
}
impl BatchQueue {
fn new() -> Self {
Self {
inner: Mutex::new(BatchQueueState {
batches: VecDeque::new(),
closed: false,
}),
cond: Condvar::new(),
}
}
fn push(&self, batch: Batch) {
let mut state = self.inner.lock().unwrap_or_else(PoisonError::into_inner);
state.batches.push_back(batch);
drop(state);
self.cond.notify_one();
}
fn close(&self) {
let mut state = self.inner.lock().unwrap_or_else(PoisonError::into_inner);
state.closed = true;
drop(state);
self.cond.notify_all();
}
fn pop(&self) -> Option<Batch> {
let mut state = self.inner.lock().unwrap_or_else(PoisonError::into_inner);
loop {
if let Some(batch) = state.batches.pop_front() {
return Some(batch);
}
if state.closed {
return None;
}
state = self
.cond
.wait(state)
.unwrap_or_else(PoisonError::into_inner);
}
}
}
struct ByteBudget {
state: Mutex<ByteBudgetState>,
cond: Condvar,
cap: u64,
}
struct ByteBudgetState {
used: u64,
shutdown: bool,
}
impl ByteBudget {
fn new(cap: u64) -> Self {
Self {
state: Mutex::new(ByteBudgetState {
used: 0,
shutdown: false,
}),
cond: Condvar::new(),
cap: cap.max(1),
}
}
fn acquire(&self, n: u64) -> bool {
let mut state = self.state.lock().unwrap_or_else(PoisonError::into_inner);
loop {
if state.shutdown {
return false;
}
if state.used == 0 || state.used + n <= self.cap {
state.used += n;
return true;
}
state = self
.cond
.wait(state)
.unwrap_or_else(PoisonError::into_inner);
}
}
fn release(&self, n: u64) {
let mut state = self.state.lock().unwrap_or_else(PoisonError::into_inner);
state.used = state.used.saturating_sub(n);
drop(state);
self.cond.notify_one();
}
fn shutdown(&self) {
let mut state = self.state.lock().unwrap_or_else(PoisonError::into_inner);
state.shutdown = true;
drop(state);
self.cond.notify_all();
}
fn is_shutdown(&self) -> bool {
self.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.shutdown
}
}
impl ElementReader<FileReader> {
pub fn from_path<P: AsRef<Path>>(path: P) -> Result<Self> {
let mut blob_iter = BlobReader::from_path(path)?;
let header = read_header_blob(&mut blob_iter)?;
Ok(ElementReader {
blob_iter,
header,
decode_threads: None,
pipeline_config: PipelineConfig::default(),
blob_filter: None,
})
}
#[cfg(feature = "linux-direct-io")]
pub fn from_path_direct<P: AsRef<Path>>(path: P) -> Result<Self> {
let mut blob_iter = BlobReader::from_path_direct(path)?;
let header = read_header_blob(&mut blob_iter)?;
Ok(ElementReader {
blob_iter,
header,
decode_threads: None,
pipeline_config: PipelineConfig::default(),
blob_filter: None,
})
}
pub fn open<P: AsRef<Path>>(path: P, direct: bool) -> Result<Self> {
let mut blob_iter = BlobReader::open(path, direct)?;
let header = read_header_blob(&mut blob_iter)?;
Ok(ElementReader {
blob_iter,
header,
decode_threads: None,
pipeline_config: PipelineConfig::default(),
blob_filter: None,
})
}
}
fn read_header_blob<R: Read + Send>(blob_iter: &mut BlobReader<R>) -> Result<HeaderBlock> {
match blob_iter.next() {
Some(Ok(blob)) => match blob.decode()? {
BlobDecode::OsmHeader(header) => Ok(*header),
_ => Err(new_error(ErrorKind::MissingHeader)),
},
Some(Err(e)) => Err(e),
None => Err(new_error(ErrorKind::MissingHeader)),
}
}
fn node_id(element: &Element<'_>) -> Option<i64> {
match element {
Element::Node(n) => Some(n.id()),
Element::DenseNode(n) => Some(n.id()),
_ => None,
}
}
pub struct PipelinedBlocks {
rx: Option<Receiver<Result<PrimitiveBlock>>>,
handle: Option<JoinHandle<()>>,
}
impl Iterator for PipelinedBlocks {
type Item = Result<PrimitiveBlock>;
fn next(&mut self) -> Option<Self::Item> {
self.rx.as_ref()?.recv().ok()
}
}
impl Drop for PipelinedBlocks {
fn drop(&mut self) {
drop(self.rx.take());
if let Some(h) = self.handle.take() {
drop(h.join());
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::{Batch, BatchQueue, ByteBudget, ElementReader, par_fold_blobs};
use crate::error::{BlobError, ErrorKind};
use std::io::{Cursor, Read};
use std::sync::Arc;
use std::sync::mpsc;
use std::time::Duration;
#[test]
fn byte_budget_blocks_at_cap_until_release() {
let budget = Arc::new(ByteBudget::new(10));
assert!(budget.acquire(6));
let (ready_tx, ready_rx) = mpsc::sync_channel(1);
let (done_tx, done_rx) = mpsc::sync_channel(1);
let child = Arc::clone(&budget);
let handle = std::thread::spawn(move || {
ready_tx.send(()).unwrap();
let admitted = child.acquire(5);
done_tx.send(admitted).unwrap();
});
ready_rx.recv().unwrap();
assert!(
done_rx.recv_timeout(Duration::from_millis(50)).is_err(),
"acquire must block while the budget is over cap"
);
budget.release(6);
assert!(
done_rx.recv_timeout(Duration::from_secs(5)).unwrap(),
"acquire must succeed once capacity frees up"
);
handle.join().unwrap();
}
#[test]
fn byte_budget_admits_oversized_when_empty() {
let budget = ByteBudget::new(10);
assert!(budget.acquire(100));
}
#[test]
fn byte_budget_shutdown_unblocks_acquirer() {
let budget = Arc::new(ByteBudget::new(10));
assert!(budget.acquire(10));
let (ready_tx, ready_rx) = mpsc::sync_channel(1);
let (done_tx, done_rx) = mpsc::sync_channel(1);
let child = Arc::clone(&budget);
let handle = std::thread::spawn(move || {
ready_tx.send(()).unwrap();
let admitted = child.acquire(5);
done_tx.send(admitted).unwrap();
});
ready_rx.recv().unwrap();
assert!(done_rx.recv_timeout(Duration::from_millis(50)).is_err());
budget.shutdown();
assert!(
!done_rx.recv_timeout(Duration::from_secs(5)).unwrap(),
"shutdown must wake a blocked acquirer and return false"
);
handle.join().unwrap();
}
#[test]
fn batch_queue_drains_then_closes() {
let queue = BatchQueue::new();
queue.push(Batch::new());
queue.close();
assert!(queue.pop().is_some(), "queued batch drains after close");
assert!(queue.pop().is_none(), "closed empty queue returns None");
}
#[test]
fn batch_queue_close_wakes_blocked_consumer() {
let queue = Arc::new(BatchQueue::new());
let (done_tx, done_rx) = mpsc::sync_channel(1);
let child = Arc::clone(&queue);
let handle = std::thread::spawn(move || {
let got = child.pop();
done_tx.send(got.is_none()).unwrap();
});
assert!(done_rx.recv_timeout(Duration::from_millis(50)).is_err());
queue.close();
assert!(
done_rx.recv_timeout(Duration::from_secs(5)).unwrap(),
"close must wake a blocked consumer with None"
);
handle.join().unwrap();
}
fn build_pbf(data_blobs: usize) -> Vec<u8> {
use crate::block_builder::{BlockBuilder, HeaderBuilder};
use crate::writer::{Compression, PbfWriter};
let mut buf = Vec::new();
{
let mut writer = PbfWriter::new(&mut buf, Compression::Zlib(6));
let header = HeaderBuilder::new().build().unwrap();
writer.write_header(&header).unwrap();
for _ in 0..data_blobs {
let mut bb = BlockBuilder::new();
for i in 1..=4_i32 {
bb.add_node(
i64::from(i),
500_000_000 + i,
100_000_000 + i,
std::iter::empty::<(&str, &str)>(),
None,
);
}
let block = bb.take().unwrap().unwrap();
writer.write_primitive_block(block).unwrap();
}
writer.flush().unwrap();
}
buf
}
fn first_data_blob_offset(buf: &[u8]) -> usize {
use crate::read::blob::BlobReader;
let mut r = BlobReader::new_seekable(Cursor::new(buf.to_vec())).unwrap();
let _header = r.next().unwrap().unwrap();
let first_data = r.next().unwrap().unwrap();
usize::try_from(first_data.offset().unwrap().0).unwrap()
}
#[allow(clippy::unwrap_in_result)]
fn assert_completes<F, T>(label: &str, f: F) -> std::thread::Result<T>
where
F: FnOnce() -> T + Send + 'static,
T: Send + 'static,
{
let (tx, rx) = mpsc::sync_channel::<()>(1);
let handle = std::thread::spawn(move || {
let out = std::panic::catch_unwind(std::panic::AssertUnwindSafe(f));
if tx.send(()).is_err() {
}
out
});
match rx.recv_timeout(Duration::from_secs(30)) {
Ok(()) => handle
.join()
.expect("watchdog thread itself failed to join"),
Err(_) => panic!("{label}: did not complete within timeout - deadlock"),
}
}
struct PanicAfterRead {
data: Vec<u8>,
pos: usize,
panic_at: usize,
}
impl Read for PanicAfterRead {
fn read(&mut self, out: &mut [u8]) -> std::io::Result<usize> {
assert!(
self.pos <= self.panic_at,
"reader served past its panic threshold"
);
if self.pos >= self.panic_at {
panic!("boom in Read during pump");
}
let end = (self.pos + out.len()).min(self.panic_at);
let n = end - self.pos;
out[..n].copy_from_slice(&self.data[self.pos..end]);
self.pos += n;
Ok(n)
}
}
#[test]
fn par_map_reduce_accepts_borrowed_non_static_reader() {
let buf = build_pbf(3);
let reader = ElementReader::new(Cursor::new(buf.as_slice())).unwrap();
let count = reader
.par_map_reduce(|_e| 1_u64, || 0_u64, |a, b| a + b)
.unwrap();
assert_eq!(count, 12, "3 blobs x 4 dense nodes");
}
#[test]
fn worker_panic_one_worker_over_budget_does_not_deadlock() {
let buf = build_pbf(4);
let result = assert_completes("worker panic", move || {
let ElementReader { blob_iter, .. } = ElementReader::new(Cursor::new(buf)).unwrap();
par_fold_blobs(
blob_iter,
1, 1, 1, super::PAR_BATCH_MAX_BYTES,
|_e| -> u64 { panic!("boom in map_op") },
|| 0_u64,
|a, b| a + b,
)
});
assert!(
result.is_err(),
"a worker panic must propagate as a panic, not be swallowed"
);
}
#[test]
fn panicking_read_during_pump_does_not_deadlock() {
let buf = build_pbf(3);
let panic_at = first_data_blob_offset(&buf);
let result = assert_completes("panicking read", move || {
let reader = ElementReader::new(PanicAfterRead {
data: buf,
pos: 0,
panic_at,
})
.unwrap();
reader.par_map_reduce(|_e| 1_u64, || 0_u64, |a, b| a + b)
});
assert!(
result.is_err(),
"a panicking Read must propagate, not deadlock scope cleanup"
);
}
#[test]
fn read_error_wins_over_simultaneous_decode_error() {
let mut buf = build_pbf(1);
let n = buf.len();
for b in &mut buf[n - 4..n] {
*b ^= 0xFF;
}
buf.extend_from_slice(&[0x00, 0x01, 0x00, 0x00]);
let result = assert_completes("read vs decode", move || {
let ElementReader { blob_iter, .. } = ElementReader::new(Cursor::new(buf)).unwrap();
par_fold_blobs(
blob_iter,
2,
super::PAR_INFLIGHT_BUDGET,
1, super::PAR_BATCH_MAX_BYTES,
|_e| 1_u64,
|| 0_u64,
|a, b| a + b,
)
})
.expect("orchestration must not panic");
let err = result.expect_err("must surface an error");
assert!(
matches!(err.kind(), ErrorKind::Blob(BlobError::HeaderTooBig { .. })),
"framing/read error must win over the decode error, got {err:?}"
);
}
}