use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncWrite, AsyncWriteExt};
use tape_rpc::Rpc;
use tape_api::program::tapedrive::track_pda;
use tape_core::track::types::CompressedTrack;
use tape_core::types::StorageUnits;
use tape_crypto::address::Address;
use tape_protocol::Api;
use crate::error::TapedriveError;
use crate::metrics::{Operation, Phase};
use crate::tapedrive::Tapedrive;
use super::error::StreamError;
use super::manifest::ChunkManifest;
pub async fn read_bytes<Blockchain: Rpc, Cluster: Api>(
client: &Tapedrive<Blockchain, Cluster>,
manifest_address: &Address,
) -> Result<Vec<u8>, TapedriveError> {
let (manifest, manifest_track) = read_manifest(client, manifest_address).await?;
let mut buffer = MemoryWriter::new(manifest.total_size)?;
read_manifest_into(client, &manifest, &manifest_track, &mut buffer).await?;
Ok(buffer.into_inner())
}
pub async fn read_into<Blockchain: Rpc, Cluster: Api, Writer: AsyncWrite + Unpin>(
client: &Tapedrive<Blockchain, Cluster>,
manifest_address: &Address,
mut writer: Writer,
) -> Result<(), TapedriveError> {
let (manifest, manifest_track) = read_manifest(client, manifest_address).await?;
read_manifest_into(client, &manifest, &manifest_track, &mut writer).await
}
pub(crate) async fn read_manifest<Blockchain: Rpc, Cluster: Api>(
client: &Tapedrive<Blockchain, Cluster>,
manifest_address: &Address,
) -> Result<(ChunkManifest, CompressedTrack), TapedriveError> {
let manifest_bytes = client
.read_as(manifest_address, Operation::ReadStream)
.await?;
let manifest = ChunkManifest::from_bytes(&manifest_bytes)
.map_err(|error| stream_error(StreamError::Manifest(format!("invalid manifest: {error}"))))?;
let metadata = client.timer(Operation::ReadStream, Phase::TrackMetadata);
let manifest_track = client.get_track(manifest_address).await;
metadata.finish_result(&manifest_track);
let manifest_track = manifest_track?;
Ok((manifest, manifest_track))
}
async fn read_manifest_into<Blockchain: Rpc, Cluster: Api, Writer: AsyncWrite + Unpin>(
client: &Tapedrive<Blockchain, Cluster>,
manifest: &ChunkManifest,
manifest_track: &CompressedTrack,
writer: &mut Writer,
) -> Result<(), TapedriveError> {
let tape_address = manifest_track.tape;
let mut total_written = StorageUnits::zero();
for (chunk_index, entry) in manifest.chunks.iter().enumerate() {
let track_address = track_pda(tape_address, entry.track_number).0;
let chunk_data = client.read_as(&track_address, Operation::ReadStream).await?;
let data_size = StorageUnits::from_bytes(chunk_data.len() as u64);
if data_size != entry.size {
return Err(stream_error(StreamError::Chunk(format!(
"chunk {chunk_index} size mismatch: expected {}, got {}",
entry.size,
chunk_data.len()
))));
}
let write_sink = client
.timer(Operation::ReadStream, Phase::WriteSink)
.bytes(chunk_data.len() as u64);
let result = writer.write_all(&chunk_data).await;
write_sink.finish_result(&result);
result?;
total_written = total_written
.checked_add(data_size)
.ok_or_else(|| stream_error(StreamError::Integrity("stream size overflow".into())))?;
}
let write_sink = client.timer(Operation::ReadStream, Phase::WriteSink);
let result = writer.flush().await;
write_sink.finish_result(&result);
result?;
if total_written != manifest.total_size {
return Err(stream_error(StreamError::Integrity(format!(
"reassembled size mismatch: expected {}, got {total_written}",
manifest.total_size
))));
}
Ok(())
}
fn stream_error(error: StreamError) -> TapedriveError {
TapedriveError::Stream(error.to_string())
}
struct MemoryWriter {
data: Vec<u8>,
}
impl MemoryWriter {
fn new(total_size: StorageUnits) -> Result<Self, TapedriveError> {
let capacity = usize::try_from(total_size.to_bytes()).map_err(|_| {
stream_error(StreamError::InvalidInput(
"stream too large to fit in memory".into(),
))
})?;
Ok(Self {
data: Vec::with_capacity(capacity),
})
}
fn into_inner(self) -> Vec<u8> {
self.data
}
}
impl AsyncWrite for MemoryWriter {
fn poll_write(
mut self: Pin<&mut Self>,
_ctx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
self.data.extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(
self: Pin<&mut Self>,
_ctx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(
self: Pin<&mut Self>,
_ctx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
}