use std::{
collections::{HashMap, HashSet, VecDeque},
sync::Arc,
};
use async_trait::async_trait;
use lunchbox::{
path::PathBuf,
types::{
DirEntry, HasFileType, MaybeSend, MaybeSync, Metadata, PathType, ReadDir, ReadDirPoller,
ReadableFile,
},
ReadableFileSystem,
};
use pin_project::pin_project;
use tokio::io::AsyncRead;
pub(crate) struct OverlayFS<B, T> {
bottom: Arc<B>,
top: Arc<T>,
}
impl<B, T> OverlayFS<B, T> {
pub fn new(bottom: Arc<B>, top: Arc<T>) -> Self {
Self { bottom, top }
}
}
#[pin_project(project = EnumProj)]
pub(crate) enum OverlayFile<B, T> {
Bottom(#[pin] B),
Top(#[pin] T),
}
impl<B: AsyncRead, T: AsyncRead> AsyncRead for OverlayFile<B, T> {
fn poll_read(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
match self.project() {
EnumProj::Bottom(v) => v.poll_read(cx, buf),
EnumProj::Top(v) => v.poll_read(cx, buf),
}
}
}
#[cfg_attr(target_family = "wasm", async_trait(?Send))]
#[cfg_attr(not(target_family = "wasm"), async_trait)]
impl<B, T> ReadableFile for OverlayFile<B, T>
where
B: ReadableFile + MaybeSend + MaybeSync,
T: ReadableFile + MaybeSend + MaybeSync,
{
async fn metadata(&self) -> std::io::Result<Metadata> {
match self {
OverlayFile::Bottom(v) => v.metadata().await,
OverlayFile::Top(v) => v.metadata().await,
}
}
async fn try_clone(&self) -> std::io::Result<Self> {
match self {
OverlayFile::Bottom(v) => v.try_clone().await.map(Self::Bottom),
OverlayFile::Top(v) => v.try_clone().await.map(Self::Top),
}
}
}
impl<B: HasFileType, T: HasFileType> HasFileType for OverlayFS<B, T> {
type FileType = OverlayFile<B::FileType, T::FileType>;
}
#[cfg_attr(target_family = "wasm", async_trait(?Send))]
#[cfg_attr(not(target_family = "wasm"), async_trait)]
impl<B, T> ReadableFileSystem for OverlayFS<B, T>
where
B: ReadableFileSystem + MaybeSend + MaybeSync,
T: ReadableFileSystem + MaybeSend + MaybeSync,
B::FileType: ReadableFile + MaybeSend + MaybeSync,
T::FileType: ReadableFile + MaybeSend + MaybeSync,
B::ReadDirPollerType: MaybeSend,
T::ReadDirPollerType: MaybeSend,
{
async fn open(&self, path: impl PathType) -> std::io::Result<Self::FileType>
where
Self::FileType: ReadableFile,
{
let p = &self.canonicalize(path).await?;
fallthrough(
|| async { self.top.open(p).await.map(OverlayFile::Top) },
|| async { self.bottom.open(p).await.map(OverlayFile::Bottom) },
)
.await
}
async fn canonicalize(&self, path: impl PathType) -> std::io::Result<PathBuf> {
let mut path: PathBuf = path.as_ref().to_owned();
let mut visited = HashSet::new();
loop {
path = path_clean::clean(path.as_str()).into();
if visited.contains(&path) {
return Err(std::io::Error::new(
std::io::ErrorKind::Other,
"Found symlink loop",
));
}
visited.insert(path.clone());
let f = self.read_link(&path).await;
if f.is_err() {
return Ok(path);
}
path = f.unwrap();
}
}
async fn metadata(&self, path: impl PathType) -> std::io::Result<Metadata> {
let p = &self.canonicalize(path).await?;
fallthrough(
|| async { self.top.metadata(p).await },
|| async { self.bottom.metadata(p).await },
)
.await
}
async fn read(&self, path: impl PathType) -> std::io::Result<Vec<u8>> {
let p = &self.canonicalize(path).await?;
fallthrough(
|| async { self.top.read(p).await },
|| async { self.bottom.read(p).await },
)
.await
}
type ReadDirPollerType = OverlayReadDirPoller;
async fn read_dir(
&self,
path: impl PathType,
) -> std::io::Result<ReadDir<Self::ReadDirPollerType, Self>> {
let p = path.as_ref();
let mut entries = HashMap::new();
let top_info = self.top.read_dir(p).await;
let bottom_info = self.bottom.read_dir(p).await;
if top_info.is_err() {
if let Err(e) = bottom_info {
return Err(e);
}
}
if let Ok(mut dir) = bottom_info {
while let Some(entry) = dir.next_entry().await? {
let p = entry.path();
entries.insert(
p.clone(),
Entry {
file_name: entry.file_name(),
path: p,
},
);
}
}
if let Ok(mut dir) = top_info {
while let Some(entry) = dir.next_entry().await? {
let p = entry.path();
entries.insert(
p.clone(),
Entry {
file_name: entry.file_name(),
path: p,
},
);
}
}
let poller = OverlayReadDirPoller {
entries: entries.into_iter().map(|(_, v)| v).collect(),
};
Ok(ReadDir::new(poller, self))
}
async fn read_link(&self, path: impl PathType) -> std::io::Result<PathBuf> {
let p = path.as_ref();
fallthrough(
|| async { self.top.read_link(p).await },
|| async { self.bottom.read_link(p).await },
)
.await
}
async fn read_to_string(&self, path: impl PathType) -> std::io::Result<String> {
let p = &self.canonicalize(path).await?;
fallthrough(
|| async { self.top.read_to_string(p).await },
|| async { self.bottom.read_to_string(p).await },
)
.await
}
async fn symlink_metadata(&self, path: impl PathType) -> std::io::Result<Metadata> {
let p = path.as_ref();
fallthrough(
|| async { self.top.symlink_metadata(p).await },
|| async { self.bottom.symlink_metadata(p).await },
)
.await
}
}
pub(crate) struct OverlayReadDirPoller {
entries: VecDeque<Entry>,
}
struct Entry {
file_name: String,
path: PathBuf,
}
impl<F> ReadDirPoller<F> for OverlayReadDirPoller
where
F: ReadableFileSystem,
F::FileType: ReadableFile,
{
fn poll_next_entry<'a>(
&mut self,
_cx: &mut std::task::Context<'_>,
fs: &'a F,
) -> std::task::Poll<std::io::Result<Option<lunchbox::types::DirEntry<'a, F>>>> {
std::task::Poll::Ready(Ok(self
.entries
.pop_front()
.map(|v| DirEntry::new(fs, v.file_name, v.path))))
}
}
async fn fallthrough<T, U, R, Fut, Fut2>(t: T, u: U) -> std::io::Result<R>
where
T: FnOnce() -> Fut,
U: FnOnce() -> Fut2,
Fut: std::future::Future<Output = std::io::Result<R>>,
Fut2: std::future::Future<Output = std::io::Result<R>>,
{
if let Ok(v) = t().await {
return Ok(v);
}
u().await
}