use crate::sampling::ReadWatch;
use color_eyre::Result;
use color_eyre::eyre::eyre;
use polars::lazy::dsl::{DslPlan, ScanSources};
use polars::prelude::{LazyFrame, PlRefPath};
use std::collections::HashMap;
use std::io::Write;
use std::path::{Path, PathBuf};
use std::sync::Arc;
pub const COPIES_DIR: &str = "quality-copies";
const HELD: &str = ".held";
const UNHELD_GRACE: std::time::Duration = std::time::Duration::from_secs(60);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RemoteObject {
pub url: String,
pub size: u64,
pub etag: Option<String>,
}
#[derive(Debug)]
pub struct LocalCopy {
_held: std::fs::File,
dir: tempfile::TempDir,
paths: HashMap<String, PathBuf>,
bytes: u64,
}
impl LocalCopy {
pub fn bytes(&self) -> u64 {
self.bytes
}
pub fn objects(&self) -> usize {
self.paths.len()
}
pub fn dir(&self) -> &Path {
self.dir.path()
}
pub fn covers(&self, url: &str) -> bool {
self.paths.contains_key(url)
}
pub fn fetch(
root: &Path,
objects: &[RemoteObject],
stop: &ReadWatch,
mut get: impl FnMut(&RemoteObject, &mut dyn FnMut(&[u8]) -> Result<()>) -> Result<()>,
) -> Result<LocalCopy> {
std::fs::create_dir_all(root).map_err(unwritable)?;
sweep(root);
let dir = tempfile::Builder::new()
.prefix("copy-")
.tempdir_in(root)
.map_err(unwritable)?;
let held = hold(dir.path())?;
let mut copy = LocalCopy {
_held: held,
dir,
paths: HashMap::with_capacity(objects.len()),
bytes: 0,
};
for object in objects {
stop.check()?;
let path = copy.dir.path().join(relative_path(&object.url));
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent).map_err(unwritable)?;
}
let mut file = std::fs::File::create_new(&path).map_err(unwritable)?;
let mut written = 0u64;
get(object, &mut |chunk: &[u8]| {
stop.check()?;
file.write_all(chunk).map_err(unwritable)?;
written += chunk.len() as u64;
Ok(())
})?;
file.flush().map_err(unwritable)?;
if written != object.size {
return Err(eyre!(
"{} is {written} bytes, not the {} listed when it opened: \
it changed. Open the dataset again",
object.url,
object.size
));
}
copy.bytes += written;
copy.paths.insert(object.url.clone(), path);
}
Ok(copy)
}
pub fn redirect(&self, lf: &LazyFrame) -> Option<LazyFrame> {
let mut plan = lf.logical_plan.clone();
if !redirect_plan(&mut plan, &self.paths) || reads_remote(&plan) {
return None;
}
Some(LazyFrame::from(plan).with_optimizations(lf.get_current_optimizations()))
}
}
fn unwritable(error: std::io::Error) -> color_eyre::Report {
eyre!(
"Could not write the local copy: {error}. \
quality_local_copy = 0 reads the source instead"
)
}
pub fn same_etag(listed: &str, fetched: &str) -> bool {
let bare = |tag: &str| {
tag.trim()
.trim_start_matches("W/")
.trim_matches('"')
.to_string()
};
bare(listed) == bare(fetched)
}
fn hold(dir: &Path) -> Result<std::fs::File> {
use fs2::FileExt;
let file = std::fs::File::create(dir.join(HELD)).map_err(unwritable)?;
file.try_lock_exclusive().map_err(unwritable)?;
Ok(file)
}
pub fn sweep(root: &Path) {
use fs2::FileExt;
let Ok(entries) = std::fs::read_dir(root) else {
return;
};
for entry in entries.flatten() {
let path = entry.path();
let ours = path.is_dir()
&& path
.file_name()
.and_then(|name| name.to_str())
.is_some_and(|name| name.starts_with("copy-"));
if !ours {
continue;
}
let old = |meta: std::io::Result<std::fs::Metadata>| {
meta.and_then(|meta| meta.modified())
.ok()
.and_then(|modified| modified.elapsed().ok())
.is_some_and(|age| age > UNHELD_GRACE)
};
let orphaned = match std::fs::File::open(path.join(HELD)) {
Ok(file) => {
let free = file.try_lock_exclusive().is_ok();
if free {
let _ = fs2::FileExt::unlock(&file);
}
free && old(file.metadata())
}
Err(_) => old(entry.metadata()),
};
if orphaned {
let _ = std::fs::remove_dir_all(&path);
}
}
}
pub fn free_space(dir: &Path) -> Option<u64> {
let existing = dir.ancestors().find(|path| path.exists())?;
fs2::available_space(existing).ok()
}
fn relative_path(url: &str) -> PathBuf {
let rest = url.split_once("://").map_or(url, |(_, rest)| rest);
rest.split('/')
.filter(|part| !part.is_empty())
.map(|part| match part {
"." | ".." => "_".to_string(),
part => safe_component(part, cfg!(windows)),
})
.collect()
}
fn safe_component(part: &str, windows: bool) -> String {
if !windows {
return part.to_string();
}
part.chars()
.map(|c| match c {
'<' | '>' | ':' | '"' | '\\' | '|' | '?' | '*' => format!("%{:02X}", c as u32),
c => c.to_string(),
})
.collect()
}
fn is_remote(path: &str) -> bool {
crate::source::is_remote_url(Path::new(path))
}
pub fn scan_paths(lf: &LazyFrame) -> Vec<String> {
let mut paths = Vec::new();
for node in &lf.logical_plan {
if let DslPlan::Scan {
sources: ScanSources::Paths(sources),
..
} = node
{
paths.extend(sources.iter().map(|path| path.as_str().to_string()));
}
}
paths
}
fn reads_remote(plan: &DslPlan) -> bool {
plan.into_iter().any(|node| match node {
DslPlan::Scan {
sources: ScanSources::Paths(sources),
..
} => sources.iter().any(|path| is_remote(path.as_str())),
_ => false,
})
}
fn redirect_plan(plan: &mut DslPlan, paths: &HashMap<String, PathBuf>) -> bool {
let into = |input: &mut Arc<DslPlan>| redirect_plan(Arc::make_mut(input), paths);
let each = |inputs: &mut [DslPlan]| inputs.iter_mut().all(|input| redirect_plan(input, paths));
match plan {
DslPlan::Scan {
sources,
unified_scan_args,
cached_ir,
..
} => {
let ScanSources::Paths(urls) = sources else {
return true;
};
if !urls.iter().any(|url| is_remote(url.as_str())) {
return true;
}
let mut local = Vec::with_capacity(urls.len());
for url in urls.iter() {
let Some(path) = paths.get(url.as_str()).and_then(|path| path.to_str()) else {
return false;
};
local.push(PlRefPath::new(path));
}
*sources = ScanSources::Paths(local.into_iter().collect());
unified_scan_args.cloud_options = None;
unified_scan_args.glob = false;
*cached_ir = Default::default();
true
}
DslPlan::IR { dsl, .. } => {
let mut inner = Arc::unwrap_or_clone(dsl.clone());
let ok = redirect_plan(&mut inner, paths);
*plan = inner;
ok
}
DslPlan::Select { input, .. }
| DslPlan::GroupBy { input, .. }
| DslPlan::Filter { input, .. }
| DslPlan::Distinct { input, .. }
| DslPlan::Sort { input, .. }
| DslPlan::Slice { input, .. }
| DslPlan::HStack { input, .. }
| DslPlan::MatchToSchema { input, .. }
| DslPlan::MapFunction { input, .. }
| DslPlan::Sink { input, .. }
| DslPlan::Cache { input, .. }
| DslPlan::Pivot { input, .. } => into(input),
DslPlan::Union { inputs, .. }
| DslPlan::HConcat { inputs, .. }
| DslPlan::SinkMultiple { inputs } => each(inputs),
DslPlan::PipeWithSchema { input, .. } => {
let mut inputs = input.to_vec();
let ok = each(&mut inputs);
*input = inputs.into();
ok
}
DslPlan::Join {
input_left,
input_right,
..
} => into(input_left) & into(input_right),
DslPlan::Gather { input, idxs, .. } => into(input) & into(idxs),
DslPlan::ExtContext { input, contexts } => into(input) & each(contexts),
_ => true,
}
}
#[cfg(test)]
mod tests {
use super::*;
use polars::prelude::*;
fn files(dir: &Path) -> Vec<RemoteObject> {
(0..2)
.map(|part| {
let mut df =
df!("id" => ((part * 10)..(part * 10 + 10)).collect::<Vec<i64>>()).unwrap();
let path = dir.join(format!("part={part}")).join("data.parquet");
std::fs::create_dir_all(path.parent().unwrap()).unwrap();
ParquetWriter::new(std::fs::File::create(&path).unwrap())
.finish(&mut df)
.unwrap();
let size = std::fs::metadata(&path).unwrap().len();
RemoteObject {
url: format!("s3://lake/events/part={part}/data.parquet"),
size,
etag: None,
}
})
.collect()
}
fn bytes_of(source: &Path, url: &str) -> Vec<u8> {
let key = url.trim_start_matches("s3://lake/events/");
std::fs::read(source.join(key)).unwrap()
}
fn remote_scan(urls: &[String]) -> LazyFrame {
let sources = ScanSources::Paths(urls.iter().map(PlRefPath::new).collect());
let args = polars::lazy::dsl::UnifiedScanArgs {
hive_options: polars::io::HiveOptions::new_enabled(),
..Default::default()
};
DslBuilder::scan_parquet(sources, Default::default(), args)
.unwrap()
.build()
.into()
}
#[test]
fn a_copy_reads_as_the_remote_scan_would() {
let source = tempfile::tempdir().unwrap();
let root = tempfile::tempdir().unwrap();
let objects = files(source.path());
let copy = LocalCopy::fetch(
root.path(),
&objects,
&ReadWatch::default(),
|object, write| {
for chunk in bytes_of(source.path(), &object.url).chunks(7) {
write(chunk)?;
}
Ok(())
},
)
.unwrap();
assert_eq!(copy.objects(), 2);
assert_eq!(
copy.bytes(),
objects.iter().map(|object| object.size).sum::<u64>()
);
let urls = objects
.iter()
.map(|object| object.url.clone())
.collect::<Vec<_>>();
let remote = remote_scan(&urls)
.filter(col("id").gt(lit(3)))
.group_by([col("part")])
.agg([len().alias("rows")])
.sort(["part"], Default::default());
let remote: LazyFrame = DslPlan::IR {
dsl: Arc::new(remote.logical_plan),
version: 0,
node: None,
opt_flags: None,
}
.into();
let local = copy.redirect(&remote).expect("every object is in the copy");
assert!(scan_paths(&local).iter().all(|path| !is_remote(path)));
let df = local.collect().unwrap();
assert_eq!(df.column("rows").unwrap().u32().unwrap().get(0), Some(6));
assert_eq!(df.column("rows").unwrap().u32().unwrap().get(1), Some(10));
assert_eq!(df.height(), 2, "the partition column survives the copy");
}
#[test]
fn a_plan_reading_an_object_not_copied_is_left_alone() {
let source = tempfile::tempdir().unwrap();
let root = tempfile::tempdir().unwrap();
let objects = files(source.path());
let copy = LocalCopy::fetch(
root.path(),
&objects[..1],
&ReadWatch::default(),
|object, write| write(&bytes_of(source.path(), &object.url)),
)
.unwrap();
let urls = objects
.iter()
.map(|object| object.url.clone())
.collect::<Vec<_>>();
assert!(copy.redirect(&remote_scan(&urls)).is_none());
}
#[test]
fn a_stopped_or_failed_fetch_leaves_no_files() {
let source = tempfile::tempdir().unwrap();
let root = tempfile::tempdir().unwrap();
let objects = files(source.path());
let entries = || std::fs::read_dir(root.path()).unwrap().count();
let stop = ReadWatch::default();
let stopped = LocalCopy::fetch(root.path(), &objects, &stop, |object, write| {
let bytes = bytes_of(source.path(), &object.url);
write(&bytes[..10])?;
stop.stop();
write(&bytes[10..])
});
assert!(stopped.is_err());
assert_eq!(entries(), 0, "the partial copy is gone");
let failed = LocalCopy::fetch(
root.path(),
&objects,
&ReadWatch::default(),
|object, write| {
if object.url.contains("part=1") {
return Err(eyre!("404"));
}
write(&bytes_of(source.path(), &object.url))
},
);
assert!(failed.is_err());
assert_eq!(entries(), 0, "the first object went with it");
let short = LocalCopy::fetch(
root.path(),
&objects,
&ReadWatch::default(),
|object, write| write(&bytes_of(source.path(), &object.url)[1..]),
);
assert!(short.unwrap_err().to_string().contains("changed"));
assert_eq!(entries(), 0);
}
#[test]
fn a_copy_lives_while_held_and_a_sweep_clears_orphans() {
let source = tempfile::tempdir().unwrap();
let root = tempfile::tempdir().unwrap();
let objects = files(source.path());
let fetch = || {
LocalCopy::fetch(
root.path(),
&objects,
&ReadWatch::default(),
|object, write| write(&bytes_of(source.path(), &object.url)),
)
.unwrap()
};
let held = fetch();
let dir = held.dir().to_path_buf();
let orphan = |name: &str, age: u64| {
let dir = root.path().join(name);
std::fs::create_dir_all(dir.join("lake")).unwrap();
let lock = std::fs::File::create(dir.join(HELD)).unwrap();
let then = std::time::SystemTime::now() - std::time::Duration::from_secs(age);
lock.set_modified(then).unwrap();
dir
};
let old = orphan("copy-old", 2 * UNHELD_GRACE.as_secs());
let new = orphan("copy-new", 0);
sweep(root.path());
assert!(dir.exists(), "a held copy stays");
assert!(!old.exists(), "an orphan goes");
assert!(new.exists(), "one too new to judge stays");
drop(held);
assert!(!dir.exists(), "dropped, the copy is removed");
}
#[test]
fn two_objects_never_share_a_file() {
let source = tempfile::tempdir().unwrap();
let root = tempfile::tempdir().unwrap();
let mut objects = files(source.path());
objects[1].url = objects[0].url.replace("s3://lake/", "s3://lake//");
let error = LocalCopy::fetch(
root.path(),
&objects,
&ReadWatch::default(),
|object, write| write(&bytes_of(source.path(), &object.url)),
)
.unwrap_err();
assert!(
error.to_string().contains("quality_local_copy = 0"),
"{error}"
);
assert_eq!(std::fs::read_dir(root.path()).unwrap().count(), 0);
}
#[test]
fn urls_map_to_their_bucket_and_key() {
assert_eq!(
relative_path("s3://lake/events/region=North/part-0.parquet"),
PathBuf::from("lake/events/region=North/part-0.parquet")
);
assert_eq!(
relative_path("gs://b/../x.parquet"),
PathBuf::from("b/_/x.parquet")
);
assert_eq!(
relative_path("s3:///../../etc/passwd"),
PathBuf::from("_/_/etc/passwd")
);
assert_eq!(safe_component("at=12:00", false), "at=12:00");
assert_eq!(safe_component("at=12:00", true), "at=12%3A00");
assert_eq!(safe_component(r"a\b", true), "a%5Cb");
}
#[test]
fn etags_compare_unquoted() {
assert!(same_etag("\"abc\"", "abc"));
assert!(same_etag("W/\"abc\"", "\"abc\""));
assert!(!same_etag("\"abc\"", "\"abd\""));
}
}