use arrow_array::RecordBatch;
use arrow_schema::ArrowError;
use geopackage_core::ident::quote;
use rusqlite::Connection;
use crate::{Error, Result};
use super::BatchSource;
use super::options::ArrowReadOptions;
pub(crate) fn database_path(conn: &Connection) -> Result<Option<std::path::PathBuf>> {
let file: String = conn.query_row(
"SELECT file FROM pragma_database_list WHERE name = 'main'",
[],
|row| row.get(0),
)?;
if file.is_empty() {
return Ok(None);
}
Ok(Some(std::path::PathBuf::from(file)))
}
pub(crate) fn dense_key_span(
conn: &Connection,
table: &str,
key: &str,
) -> Result<Option<(i64, i64)>> {
let sql = format!(
"SELECT min({key}), max({key}), count(*) FROM {table}",
key = quote(key)?,
table = quote(table)?
);
let (min, max, count): (Option<i64>, Option<i64>, i64) =
conn.query_row(&sql, [], |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)))?;
let (Some(min), Some(max)) = (min, max) else {
return Ok(None);
};
let span = max.checked_sub(min).and_then(|d| d.checked_add(1));
if span != Some(count) {
return Ok(None);
}
Ok(Some((min, max)))
}
pub(crate) struct ParallelBatches {
receivers: Vec<std::sync::mpsc::Receiver<std::result::Result<WorkerMessage, ArrowError>>>,
workers: Vec<std::thread::JoinHandle<()>>,
turn: usize,
done: bool,
}
impl ParallelBatches {
pub(crate) fn spawn(
path: std::path::PathBuf,
table: String,
conversion: crate::ConversionOptions,
(first, last): (i64, i64),
batch_size: usize,
max_batch_bytes: usize,
threads: usize,
) -> Self {
let mut receivers = Vec::with_capacity(threads);
let mut workers = Vec::with_capacity(threads);
for worker in 0..threads {
let (tx, rx) = std::sync::mpsc::sync_channel(1);
let path = path.clone();
let table = table.clone();
let handle = std::thread::spawn(move || {
run_worker(
&path,
&table,
conversion,
first,
last,
batch_size,
max_batch_bytes,
threads,
worker,
&tx,
);
});
receivers.push(rx);
workers.push(handle);
}
Self {
receivers,
workers,
turn: 0,
done: false,
}
}
}
impl Iterator for ParallelBatches {
type Item = std::result::Result<RecordBatch, ArrowError>;
fn next(&mut self) -> Option<Self::Item> {
if self.done {
return None;
}
loop {
match self.receivers.get(self.turn)?.recv() {
Ok(Ok(WorkerMessage::Batch(batch))) => return Some(Ok(batch)),
Ok(Ok(WorkerMessage::WindowEnd)) => {
self.turn = (self.turn + 1) % self.receivers.len().max(1);
}
Ok(Err(error)) => {
self.done = true;
return Some(Err(error));
}
Err(_) => {
self.done = true;
return None;
}
}
}
}
}
impl Drop for ParallelBatches {
fn drop(&mut self) {
self.receivers.clear();
for worker in self.workers.drain(..) {
drop(worker.join());
}
}
}
#[expect(
clippy::too_many_arguments,
reason = "a worker's whole context, passed once at spawn; a struct would be used by this call site alone"
)]
fn run_worker(
path: &std::path::Path,
table: &str,
conversion: crate::ConversionOptions,
first: i64,
last: i64,
batch_size: usize,
max_batch_bytes: usize,
threads: usize,
worker: usize,
tx: &std::sync::mpsc::SyncSender<std::result::Result<WorkerMessage, ArrowError>>,
) {
let send_error = |error: Error| {
drop(tx.send(Err(ArrowError::ExternalError(Box::new(error)))));
};
let gpkg = match crate::GeoPackage::open_read_only(path) {
Ok(gpkg) => gpkg,
Err(error) => return send_error(error),
};
let layer = match gpkg.layer(table) {
Ok(layer) => layer.with_conversion_options(conversion),
Err(error) => return send_error(error),
};
let options = ArrowReadOptions::with_batch_size(batch_size)
.with_threads(1)
.with_max_batch_bytes(max_batch_bytes);
let mut batches = match layer.read_arrow(options) {
Ok(batches) => batches,
Err(error) => return send_error(error),
};
let BatchSource::Sequential(source) = &mut batches.source else {
return;
};
let stride = match i64::try_from(batch_size.saturating_mul(threads)) {
Ok(stride) if stride > 0 => stride,
_ => return,
};
let start = match i64::try_from(batch_size.saturating_mul(worker)) {
Ok(offset) => match first.checked_add(offset) {
Some(start) => start,
None => return,
},
Err(_) => return,
};
let mut key = start;
while key <= last {
let mut remaining = batch_size;
let mut at = key;
while remaining > 0 {
match source.read_batch_at(at, remaining) {
Ok(Some(batch)) => {
let rows = source.last_batch_rows;
if tx.send(Ok(WorkerMessage::Batch(batch))).is_err() {
return; }
if rows == 0 {
break;
}
remaining -= rows.min(remaining);
at = source.next_key;
}
Ok(None) => return,
Err(error) => return send_error(error),
}
}
if tx.send(Ok(WorkerMessage::WindowEnd)).is_err() {
return; }
match key.checked_add(stride) {
Some(next) => key = next,
None => return,
}
}
}
enum WorkerMessage {
Batch(RecordBatch),
WindowEnd,
}