snix_store/nar/renderer/seekable/
mod.rs1use std::{
2 collections::HashMap,
3 io::{self, SeekFrom},
4 pin::Pin,
5 task::{Context, Poll},
6};
7
8use futures::ready;
9use pin_project::pin_project;
10use segments::Segments;
11use snix_castore::blobservice::BlobService;
12use snix_castore::directoryservice::DirectoryService;
13use snix_castore::{Node, directoryservice::DirectoryServiceGraphExt};
14use tokio::io::{AsyncBufRead, AsyncRead, AsyncSeek, AsyncWrite};
15use tracing::{instrument, warn};
16
17use crate::nar::RenderError;
18
19mod segments;
20
21#[cfg(test)]
22mod test;
23
24const SEGMENT_CONCURRENCY: usize = 24;
26
27pub async fn write_nar<W, BS, DS>(
28 mut w: W,
29 root_node: &Node,
30 blob_service: &BS,
31 directory_service: &DS,
32) -> Result<(), RenderError>
33where
34 W: AsyncWrite + Unpin + Send,
35 BS: BlobService,
36 DS: DirectoryService,
37{
38 let mut reader = Reader::new(root_node, blob_service, directory_service).await?;
39 tokio::io::copy_buf(&mut reader, &mut w)
40 .await
41 .map_err(RenderError::BlobService)?;
43
44 Ok(())
45}
46
47#[pin_project]
48pub struct Reader<'bs, BS: BlobService + 'bs> {
49 segments: Segments,
50 pos: u64,
51 blob_service: BS,
52 #[pin]
53 rd: Box<dyn AsyncBufRead + Send + Unpin + 'bs>,
54}
55
56impl<'bs, BS: BlobService + Clone + 'bs> Reader<'bs, BS> {
57 #[instrument(skip(blob_service, directory_service), err)]
66 pub async fn new(
67 root_node: &Node,
68 blob_service: BS,
69 directory_service: impl DirectoryService,
70 ) -> Result<Self, RenderError> {
72 let directories = if let Node::Directory { digest, .. } = root_node {
73 let directory_graph = directory_service.get_directory_graph(digest).await.map_err(RenderError::DirectoryService)?.ok_or_else(|| {
75 let err = RenderError::DirectoryNotFound(*digest, "root".into());
79 warn!(%err, "tried to render NAR, but DirectoryService didn't contain the root directory");
80 err
81 })?;
82
83 HashMap::from_iter(
84 directory_graph
85 .drain_leaves_to_root()
87 .map(|d| (d.digest(), d)),
88 )
89 } else {
90 Default::default()
92 };
93
94 let segments = Segments::from_root_node_and_directories(root_node, &directories);
95 let rd = segments.reader_for_offset(0, SEGMENT_CONCURRENCY, blob_service.clone());
96
97 Ok(Self {
98 segments,
99 pos: 0,
100 blob_service,
101 rd,
102 })
103 }
104
105 pub fn nar_size(&self) -> u64 {
106 self.segments.total_len()
107 }
108}
109
110impl<'bs, BS: BlobService> AsyncRead for Reader<'bs, BS> {
111 fn poll_read(
112 self: Pin<&mut Self>,
113 cx: &mut Context,
114 buf: &mut tokio::io::ReadBuf,
115 ) -> Poll<io::Result<()>> {
116 let this = self.project();
117
118 let bytes_read = {
119 let filled = buf.filled().len();
120 ready!(this.rd.poll_read(cx, buf))?;
121 buf.filled().len() - filled
122 };
123 *this.pos = this
124 .pos
125 .checked_add(bytes_read as u64)
126 .ok_or(std::io::Error::new(
127 std::io::ErrorKind::OutOfMemory,
128 "position > u64::MAX bytes",
129 ))?;
130
131 Poll::Ready(Ok(()))
132 }
133}
134
135impl<'bs, BS: BlobService> AsyncBufRead for Reader<'bs, BS> {
136 fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<&[u8]>> {
137 let this = self.project();
138 this.rd.poll_fill_buf(cx)
139 }
140
141 fn consume(self: Pin<&mut Self>, amt: usize) {
142 let this = self.project();
143
144 this.rd.consume(amt);
145 *this.pos = this
146 .pos
147 .checked_add(amt as u64)
148 .expect("consume would increase pos > u64::MAX bytes");
149 }
150}
151
152impl<'bs, BS: BlobService + Clone + 'bs> AsyncSeek for Reader<'bs, BS> {
153 fn start_seek(self: Pin<&mut Self>, pos: io::SeekFrom) -> io::Result<()> {
154 let nar_size = self.nar_size();
155 let new_pos = calc_pos(self.pos, nar_size, pos)?;
156
157 if new_pos != self.pos {
158 let mut this = self.project();
160
161 *this.rd = this.segments.reader_for_offset(
162 new_pos,
163 SEGMENT_CONCURRENCY,
164 this.blob_service.clone(),
165 );
166 *this.pos = new_pos;
167 }
168
169 Ok(())
170 }
171 fn poll_complete(self: Pin<&mut Self>, _cx: &mut Context) -> Poll<io::Result<u64>> {
172 Poll::Ready(Ok(self.pos))
173 }
174}
175
176fn calc_pos(cur_pos: u64, nar_size: u64, seek_from: SeekFrom) -> std::io::Result<u64> {
178 let new_pos = match seek_from {
179 SeekFrom::Start(p) => p,
180 SeekFrom::End(p) => nar_size.checked_sub_signed(p).ok_or(std::io::Error::new(
181 std::io::ErrorKind::InvalidInput,
182 "tried to seek before beginning of NAR",
183 ))?,
184 SeekFrom::Current(p) => cur_pos.checked_add_signed(p).ok_or(std::io::Error::new(
185 std::io::ErrorKind::UnexpectedEof,
186 "tried to seek way past end of NAR",
187 ))?,
188 };
189
190 if new_pos > nar_size {
191 Err(std::io::Error::new(
192 std::io::ErrorKind::UnexpectedEof,
193 "tried to seek past end of NAR",
194 ))
195 } else {
196 Ok(new_pos)
197 }
198}