use super::blob::{Blob, BlobDecode, BlobReader, BlobType};
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::{new_error, ErrorKind, Result};
use rayon::prelude::*;
use std::io::Read;
use std::path::Path;
use std::sync::mpsc::{Receiver, sync_channel};
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;
for blob in blob_iter {
match blob?.decode() {
Ok(BlobDecode::OsmData(block)) => {
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(_) => {} Err(e) => return Err(e),
}
}
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 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!(
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 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 result = super::pipeline::run_pipeline(
blob_iter,
decode_threads,
pipeline_config,
blob_filter,
|block| {
tx.send(Ok(block)).map_err(|_| {
new_error(ErrorKind::Io(std::io::Error::other(
"pipeline consumer dropped",
)))
})
},
);
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 blobs = collect_osm_data_blobs(self.blob_iter)?;
blobs
.into_par_iter()
.try_fold(
&identity,
|acc, blob: Blob| match blob.decode()? {
BlobDecode::OsmData(block) => {
Ok(block.elements().map(&map_op).fold(acc, &reduce_op))
}
BlobDecode::OsmHeader(_) | BlobDecode::Unknown(_) => Ok(acc),
},
)
.try_reduce(&identity, |a, b| Ok(reduce_op(a, b)))
}
}
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,
}
}
#[hotpath::measure]
fn collect_osm_data_blobs<R: Read + Send>(blob_iter: BlobReader<R>) -> Result<Vec<Blob>> {
let mut blobs = Vec::new();
for blob_result in blob_iter {
let blob = blob_result?;
if blob.get_type() == BlobType::OsmData {
blobs.push(blob);
}
}
Ok(blobs)
}
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());
}
}
}