diff --git a/dir-structure-tools/src/deferred_read_or_own.rs b/dir-structure-tools/src/deferred_read_or_own.rs index 6c50a46..6525ebb 100644 --- a/dir-structure-tools/src/deferred_read_or_own.rs +++ b/dir-structure-tools/src/deferred_read_or_own.rs @@ -477,6 +477,125 @@ where } } +#[cfg(feature = "async")] +#[cfg_attr(docsrs, doc(cfg(feature = "async")))] +#[pin_project(project_replace = DeferredReadOrOwnWriteRefFutureProj)] +#[doc(hidden)] +pub enum DeferredReadOrOwnWriteRefFuture< + 'r, + 'a, + T, + Vfs: WriteSupportingVfsAsync + 'r, + const CHECK_ON_READ: bool, +> where + 'r: 'a, + T: WriteToAsyncRef<'r, Vfs> + Send + 'static, + T: for<'b> ReadFromAsync<'b, Vfs> + for<'b> WriteToAsync<'b, Vfs> + Send + 'static, + for<'b> >::Future: Future> + Unpin + 'b, + for<'b> >::Future: Future> + Unpin + 'b, +{ + Poisson, + Own { + inner: >::Future<'a>, + }, + Deferred { + inner: as WriteToAsyncRef<'r, Vfs>>::Future<'a>, + }, +} + +// needed to avoid ICE in rustdoc, see https://github.com/rust-lang/rust/issues/144918 +#[cfg(all(feature = "async", doc))] +impl<'r, 'a, T, Vfs, const CHECK_ON_READ: bool> core::marker::Unpin + for DeferredReadOrOwnWriteRefFutureProj<'r, 'a, T, Vfs, CHECK_ON_READ> +where + 'r: 'a, + T: WriteToAsyncRef<'r, Vfs> + Send + 'static, + T: for<'b> ReadFromAsync<'b, Vfs> + for<'b> WriteToAsync<'b, Vfs> + Send + 'static, + for<'b> >::Future: Future> + Unpin + 'b, + for<'b> >::Future: Future> + Unpin + 'b, + Vfs: WriteSupportingVfsAsync + 'r, +{ +} + +#[cfg(feature = "async")] +#[cfg_attr(docsrs, doc(cfg(feature = "async")))] +impl<'r, 'a, const CHECK_ON_READ: bool, T, Vfs: WriteSupportingVfsAsync + 'r> Future + for DeferredReadOrOwnWriteRefFuture<'r, 'a, T, Vfs, CHECK_ON_READ> +where + 'r: 'a, + T: WriteToAsyncRef<'r, Vfs> + Send + 'static, + T: for<'b> ReadFromAsync<'b, Vfs> + for<'b> WriteToAsync<'b, Vfs> + Send + 'static, + for<'b> >::Future: Future> + Unpin + 'b, + for<'b> >::Future: Future> + Unpin + 'b, +{ + type Output = VfsResult<(), Vfs>; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + let this = self.as_mut().project_replace(Self::Poisson); + match this { + DeferredReadOrOwnWriteRefFutureProj::Own { mut inner } => { + match Pin::new(&mut inner).poll(cx) { + Poll::Ready(v) => Poll::Ready(v), + Poll::Pending => { + self.project_replace(Self::Own { inner }); + Poll::Pending + } + } + } + DeferredReadOrOwnWriteRefFutureProj::Deferred { mut inner } => { + match Pin::new(&mut inner).poll(cx) { + Poll::Ready(v) => Poll::Ready(v), + Poll::Pending => { + self.project_replace(Self::Deferred { inner }); + Poll::Pending + } + } + } + DeferredReadOrOwnWriteRefFutureProj::Poisson => { + panic!( + "DeferredReadOrOwnWriteRefFuture is in an invalid state. This is a bug in the code." + ); + } + } + } +} + +#[cfg(feature = "async")] +#[cfg_attr(docsrs, doc(cfg(feature = "async")))] +impl<'r, const CHECK_ON_READ: bool, T, Vfs: WriteSupportingVfsAsync + 'r> WriteToAsyncRef<'r, Vfs> + for DeferredReadOrOwn<'r, T, Vfs, CHECK_ON_READ> +where + T: WriteToAsyncRef<'r, Vfs> + Send + 'static, + T: for<'b> ReadFromAsync<'b, Vfs> + for<'b> WriteToAsync<'b, Vfs> + Send + 'static, + for<'b> >::Future: Future> + Unpin + 'b, + for<'b> >::Future: Future> + Unpin + 'b, +{ + type Future<'a> + = DeferredReadOrOwnWriteRefFuture<'r, 'a, T, Vfs, CHECK_ON_READ> + where + Self: 'a, + 'r: 'a, + Vfs: 'a; + + fn write_to_async_ref<'a>( + &'a self, + path: <::Path as PathType>::OwnedPath, + vfs: Pin<&'a Vfs>, + ) -> Self::Future<'a> + where + 'r: 'a, + { + match self { + DeferredReadOrOwn::Own(own) => DeferredReadOrOwnWriteRefFuture::Own { + inner: own.write_to_async_ref(path, vfs), + }, + DeferredReadOrOwn::Deferred(d) => DeferredReadOrOwnWriteRefFuture::Deferred { + inner: d.write_to_async_ref(path, vfs), + }, + } + } +} + #[cfg(feature = "resolve-path")] #[cfg_attr(docsrs, doc(cfg(feature = "resolve-path")))] impl<'a, const CHECK_ON_READ: bool, const NAME: [char; HAS_FIELD_MAX_LEN], T, Vfs> HasField diff --git a/dir-structure-tools/tests/async_tests.rs b/dir-structure-tools/tests/async_tests.rs index e81c624..1029ea4 100644 --- a/dir-structure-tools/tests/async_tests.rs +++ b/dir-structure-tools/tests/async_tests.rs @@ -89,6 +89,38 @@ async fn deferred_read() { assert_eq!(dir.f.perform_read_async().await.unwrap(), "f1"); } +/// `DeferredReadOrOwn` writes through `WriteToAsyncRef` in both `Deferred` and `Own` states. +#[tokio::test] +async fn deferred_read_or_own_write_to_async_ref() { + use dir_structure_tools::deferred_read::DeferredRead; + use dir_structure_tools::deferred_read_or_own::DeferredReadOrOwn; + use std::fs; + + let p = test_dir("deferred_read_or_own_write_to_async_ref"); + let d = p.join("dir"); + fs::create_dir_all(&d).unwrap(); + fs::write(d.join("f1.txt"), "f1").unwrap(); + + let vfs = TokioFsVfs; + + let deferred = DeferredReadOrOwn::::Deferred( + DeferredRead::read_from_async(d.join("f1.txt"), Pin::new(&vfs)) + .await + .unwrap(), + ); + deferred + .write_to_async_ref(d.join("out1.txt"), Pin::new(&vfs)) + .await + .unwrap(); + assert_eq!(fs::read_to_string(d.join("out1.txt")).unwrap(), "f1"); + + DeferredReadOrOwn::::Own("f2".to_owned()) + .write_to_async_ref(d.join("out2.txt"), Pin::new(&vfs)) + .await + .unwrap(); + assert_eq!(fs::read_to_string(d.join("out2.txt")).unwrap(), "f2"); +} + #[tokio::test] async fn read_all_directory_files() { #[derive(dir_structure::DirStructureAsync)]