Skip to main content

nominal_streaming/
upload.rs

1use std::io::Read;
2use std::io::Seek;
3use std::path::Path;
4use std::path::PathBuf;
5use std::pin::Pin;
6use std::sync::Arc;
7
8use conjure_error::Error;
9use conjure_http::client::AsyncWriteBody;
10use conjure_http::private::Stream;
11use conjure_object::BearerToken;
12use conjure_object::ResourceIdentifier;
13use conjure_object::SafeLong;
14use conjure_runtime_rustls_platform_verifier::BodyWriter;
15use conjure_runtime_rustls_platform_verifier::Client;
16use futures::StreamExt;
17use nominal_api::clients::ingest::api::AsyncIngestService;
18use nominal_api::clients::ingest::api::AsyncIngestServiceClient;
19use nominal_api::clients::upload::api::AsyncUploadService;
20use nominal_api::clients::upload::api::AsyncUploadServiceClient;
21use nominal_api::objects::api::rids::WorkspaceRid;
22use nominal_api::objects::ingest::api::AvroStreamOpts;
23use nominal_api::objects::ingest::api::CompleteMultipartUploadResponse;
24use nominal_api::objects::ingest::api::DatasetIngestTarget;
25use nominal_api::objects::ingest::api::ExistingDatasetIngestDestination;
26use nominal_api::objects::ingest::api::IngestOptions;
27use nominal_api::objects::ingest::api::IngestRequest;
28use nominal_api::objects::ingest::api::IngestResponse;
29use nominal_api::objects::ingest::api::IngestSource;
30use nominal_api::objects::ingest::api::InitiateMultipartUploadRequest;
31use nominal_api::objects::ingest::api::InitiateMultipartUploadResponse;
32use nominal_api::objects::ingest::api::Part;
33use nominal_api::objects::ingest::api::S3IngestSource;
34use tokio::sync::Semaphore;
35use tracing::error;
36use tracing::info;
37
38use crate::client::NominalApiClients;
39use crate::types::AuthProvider;
40
41const SMALL_FILE_SIZE_LIMIT: u64 = 512 * 1024 * 1024; // 512 MB
42
43#[derive(Clone)]
44pub struct AvroIngestManager {
45    pub upload_queue: async_channel::Receiver<PathBuf>,
46}
47
48impl AvroIngestManager {
49    pub fn new(
50        clients: NominalApiClients,
51        http_client: reqwest::Client,
52        handle: tokio::runtime::Handle,
53        opts: UploaderOpts,
54        upload_queue: async_channel::Receiver<PathBuf>,
55        auth_provider: impl AuthProvider + 'static,
56        data_source_rid: ResourceIdentifier,
57    ) -> Self {
58        let uploader = FileObjectStoreUploader::new(
59            clients.upload,
60            clients.ingest,
61            http_client,
62            handle.clone(),
63            opts,
64        );
65
66        let upload_queue_clone = upload_queue.clone();
67
68        handle.spawn(async move {
69            Self::run(upload_queue_clone, uploader, auth_provider, data_source_rid).await;
70        });
71
72        AvroIngestManager { upload_queue }
73    }
74
75    pub async fn run(
76        upload_queue: async_channel::Receiver<PathBuf>,
77        uploader: FileObjectStoreUploader,
78        auth_provider: impl AuthProvider + 'static,
79        data_source_rid: ResourceIdentifier,
80    ) {
81        while let Ok(file_path) = upload_queue.recv().await {
82            let file_name = file_path.to_str().unwrap_or("nmstream_file");
83            let file = std::fs::File::open(&file_path);
84            let Some(token) = auth_provider.token() else {
85                error!("Missing token for upload");
86                continue;
87            };
88            match file {
89                Ok(f) => {
90                    match upload_and_ingest_file(
91                        uploader.clone(),
92                        &token,
93                        auth_provider.workspace_rid(),
94                        f,
95                        file_name,
96                        &file_path,
97                        data_source_rid.clone(),
98                    )
99                    .await
100                    {
101                        Ok(()) => {}
102                        Err(e) => {
103                            error!(
104                                "Error uploading and ingesting file {}: {}",
105                                file_path.display(),
106                                e
107                            );
108                        }
109                    }
110                }
111                Err(e) => {
112                    error!("Failed to open file {}: {:?}", file_path.display(), e);
113                }
114            }
115        }
116    }
117}
118
119async fn upload_and_ingest_file(
120    uploader: FileObjectStoreUploader,
121    token: &BearerToken,
122    workspace_rid: Option<WorkspaceRid>,
123    file: std::fs::File,
124    file_name: &str,
125    file_path: &PathBuf,
126    data_source_rid: ResourceIdentifier,
127) -> Result<(), String> {
128    match uploader.upload(token, file, file_name, workspace_rid).await {
129        Ok(response) => {
130            match uploader
131                .ingest_avro(token, &response, data_source_rid)
132                .await
133            {
134                Ok(ingest_response) => {
135                    info!(
136                        "Successfully uploaded and ingested file {}: {:?}",
137                        file_name, ingest_response
138                    );
139                    if let Err(e) = std::fs::remove_file(file_path) {
140                        Err(format!(
141                            "Failed to remove file {}: {:?}",
142                            file_path.display(),
143                            e
144                        ))
145                    } else {
146                        info!("Removed file {}", file_path.display());
147                        Ok(())
148                    }
149                }
150                Err(e) => Err(format!("Failed to ingest file {file_name}: {e:?}")),
151            }
152        }
153        Err(e) => Err(format!("Failed to upload file {file_name}: {e:?}")),
154    }
155}
156
157#[derive(Debug, thiserror::Error)]
158pub enum UploaderError {
159    #[error("Conjure error: {0}")]
160    Conjure(String),
161    #[error("Failed to initiate multipart upload: {0}")]
162    IOError(#[from] std::io::Error),
163    #[error("Failed to upload part: {0}")]
164    HTTPError(#[from] reqwest::Error),
165    #[error("Error executing upload tasks: {0}")]
166    TokioError(#[from] tokio::task::JoinError),
167    #[error("Error: {0}")]
168    Other(String),
169}
170
171#[derive(Debug, Clone)]
172pub struct UploaderOpts {
173    pub chunk_size: usize,
174    pub max_retries: usize,
175    pub max_concurrent_uploads: usize,
176}
177
178impl Default for UploaderOpts {
179    fn default() -> Self {
180        UploaderOpts {
181            chunk_size: 512 * 1024 * 1024, // 512 MB
182            max_retries: 3,
183            max_concurrent_uploads: 1,
184        }
185    }
186}
187
188pub struct FileWriteBody {
189    file: std::fs::File,
190}
191
192impl FileWriteBody {
193    pub fn new(file: std::fs::File) -> Self {
194        FileWriteBody { file }
195    }
196}
197
198impl AsyncWriteBody<BodyWriter> for FileWriteBody {
199    async fn write_body(self: Pin<&mut Self>, w: Pin<&mut BodyWriter>) -> Result<(), Error> {
200        let mut file = self
201            .file
202            .try_clone()
203            .map_err(|e| Error::internal_safe(format!("Failed to clone file for upload: {e}")))?;
204
205        let mut buffer = Vec::new();
206        file.read_to_end(&mut buffer)
207            .map_err(|e| Error::internal_safe(format!("Failed to read bytes from file: {e}")))?;
208
209        w.write_bytes(buffer.into())
210            .await
211            .map_err(|e| Error::internal_safe(format!("Failed to write bytes to body: {e}")))?;
212
213        Ok(())
214    }
215
216    async fn reset(self: Pin<&mut Self>) -> bool {
217        let Ok(mut file) = self.file.try_clone() else {
218            return false;
219        };
220
221        use std::io::SeekFrom;
222
223        file.seek(SeekFrom::Start(0)).is_ok()
224    }
225}
226
227#[derive(Clone)]
228pub struct FileObjectStoreUploader {
229    upload_client: AsyncUploadServiceClient<Client>,
230    ingest_client: AsyncIngestServiceClient<Client>,
231    http_client: reqwest::Client,
232    handle: tokio::runtime::Handle,
233    opts: UploaderOpts,
234}
235
236impl FileObjectStoreUploader {
237    pub fn new(
238        upload_client: AsyncUploadServiceClient<Client>,
239        ingest_client: AsyncIngestServiceClient<Client>,
240        http_client: reqwest::Client,
241        handle: tokio::runtime::Handle,
242        opts: UploaderOpts,
243    ) -> Self {
244        FileObjectStoreUploader {
245            upload_client,
246            ingest_client,
247            http_client,
248            handle,
249            opts,
250        }
251    }
252
253    pub async fn initiate_upload(
254        &self,
255        token: &BearerToken,
256        file_name: &str,
257        workspace_rid: Option<WorkspaceRid>,
258    ) -> Result<InitiateMultipartUploadResponse, UploaderError> {
259        let request = InitiateMultipartUploadRequest::builder()
260            .filename(file_name)
261            .filetype("application/octet-stream")
262            .workspace(workspace_rid)
263            .build();
264        let response = self
265            .upload_client
266            .initiate_multipart_upload(token, &request)
267            .await
268            .map_err(|e| UploaderError::Conjure(format!("{e:?}")))?;
269
270        info!("Initiated multipart upload for file: {}", file_name);
271        Ok(response)
272    }
273
274    #[expect(clippy::too_many_arguments)]
275    async fn upload_part(
276        client: AsyncUploadServiceClient<Client>,
277        http_client: reqwest::Client,
278        token: BearerToken,
279        upload_id: String,
280        key: String,
281        part_number: i32,
282        chunk: Vec<u8>,
283        max_retries: usize,
284    ) -> Result<Part, UploaderError> {
285        let mut attempts = 0;
286
287        loop {
288            attempts += 1;
289            match Self::try_upload_part(
290                client.clone(),
291                http_client.clone(),
292                &token,
293                &upload_id,
294                &key,
295                part_number,
296                chunk.clone(),
297            )
298            .await
299            {
300                Ok(part) => return Ok(part),
301                Err(e) if attempts < max_retries => {
302                    error!("Upload attempt {} failed, retrying: {}", attempts, e);
303                    continue;
304                }
305                Err(e) => {
306                    return Err(e);
307                }
308            }
309        }
310    }
311
312    async fn try_upload_part(
313        client: AsyncUploadServiceClient<Client>,
314        http_client: reqwest::Client,
315        token: &BearerToken,
316        upload_id: &str,
317        key: &str,
318        part_number: i32,
319        chunk: Vec<u8>,
320    ) -> Result<Part, UploaderError> {
321        // `bucket` is None: we never set a `destination` on the initiate request, so the server
322        // opens the upload in (and resolves these follow-up calls against) the default uploads
323        // bucket. Thread `initiate_response.bucket()` through if we ever request FILE_STORE.
324        let response = client
325            .sign_part(token, upload_id, key, part_number, None)
326            .await
327            .map_err(|e| UploaderError::Conjure(format!("{e:?}")))?;
328
329        let mut request_builder = http_client.put(response.url()).body(chunk);
330
331        for (header_name, header_value) in response.headers() {
332            request_builder = request_builder.header(header_name, header_value);
333        }
334
335        let http_response = request_builder.send().await?;
336        let headers = http_response.headers().clone();
337        let status = http_response.status();
338
339        if !status.is_success() {
340            error!("Failed to upload body");
341            return Err(UploaderError::Other(format!(
342                "Failed to upload part {part_number}: HTTP status {status}"
343            )));
344        }
345
346        let etag = headers
347            .get("etag")
348            .and_then(|v| v.to_str().ok())
349            .unwrap_or("ignored-etag");
350
351        Ok(Part::new(part_number, etag))
352    }
353
354    pub async fn upload_parts<R>(
355        &self,
356        token: &BearerToken,
357        reader: R,
358        key: &str,
359        upload_id: &str,
360    ) -> Result<CompleteMultipartUploadResponse, UploaderError>
361    where
362        R: Read + Send + 'static,
363    {
364        let chunks = ChunkedStreamReader::new(reader, self.opts.chunk_size);
365
366        let parallel_part_uploads = Arc::new(Semaphore::new(self.opts.max_concurrent_uploads));
367        let mut upload_futures = Vec::new();
368
369        futures::pin_mut!(chunks);
370
371        while let Some(entry) = chunks.next().await {
372            let (index, chunk) = entry?;
373            let part_number = (index + 1) as i32;
374
375            let token = token.clone();
376            let key = key.to_string();
377            let upload_id = upload_id.to_string();
378            let parallel_part_uploads = Arc::clone(&parallel_part_uploads);
379            let client = self.upload_client.clone();
380            let http_client = self.http_client.clone();
381            let max_retries = self.opts.max_retries;
382
383            upload_futures.push(self.handle.spawn(async move {
384                let _permit = parallel_part_uploads.acquire().await;
385                Self::upload_part(
386                    client,
387                    http_client,
388                    token,
389                    upload_id,
390                    key,
391                    part_number,
392                    chunk,
393                    max_retries,
394                )
395                .await
396            }));
397        }
398
399        let mut part_responses = futures::future::join_all(upload_futures)
400            .await
401            .into_iter()
402            .map(|result| result.map_err(UploaderError::TokioError)?)
403            .collect::<Result<Vec<_>, _>>()?;
404
405        part_responses.sort_by_key(|part| part.part_number());
406
407        let response = self
408            .upload_client
409            .complete_multipart_upload(token, upload_id, key, None, &part_responses)
410            .await
411            .map_err(|e| UploaderError::Conjure(format!("{e:?}")))?;
412
413        Ok(response)
414    }
415
416    pub async fn upload_small_file(
417        &self,
418        token: &BearerToken,
419        file_name: &str,
420        size_bytes: i64,
421        workspace_rid: Option<WorkspaceRid>,
422        file: std::fs::File,
423    ) -> Result<String, UploaderError> {
424        let s3_path = self
425            .upload_client
426            .upload_file(
427                token,
428                file_name,
429                SafeLong::new(size_bytes).ok(),
430                workspace_rid.as_ref(),
431                FileWriteBody::new(file),
432            )
433            .await
434            .map_err(|e| UploaderError::Conjure(format!("{e:?}")))?;
435
436        Ok(s3_path.as_str().to_string())
437    }
438
439    pub async fn upload<R>(
440        &self,
441        token: &BearerToken,
442        reader: R,
443        file_name: impl Into<&str>,
444        workspace_rid: Option<WorkspaceRid>,
445    ) -> Result<String, UploaderError>
446    where
447        R: Read + Send + 'static,
448    {
449        let file_name = file_name.into();
450        let path = Path::new(file_name);
451        let file_size = std::fs::metadata(path)?.len();
452        if file_size < SMALL_FILE_SIZE_LIMIT {
453            return self
454                .upload_small_file(
455                    token,
456                    file_name,
457                    file_size as i64,
458                    workspace_rid,
459                    std::fs::File::open(path)?,
460                )
461                .await;
462        }
463
464        let initiate_response = self
465            .initiate_upload(token, file_name, workspace_rid)
466            .await?;
467        let upload_id = initiate_response.upload_id();
468        let key = initiate_response.key();
469
470        let response = self.upload_parts(token, reader, key, upload_id).await?;
471
472        let s3_path = response.location().ok_or_else(|| {
473            UploaderError::Other("Upload response did not contain a location".to_string())
474        })?;
475
476        Ok(s3_path.to_string())
477    }
478
479    pub async fn ingest_avro(
480        &self,
481        token: &BearerToken,
482        s3_path: &str,
483        data_source_rid: ResourceIdentifier,
484    ) -> Result<IngestResponse, UploaderError> {
485        let opts = IngestOptions::AvroStream(
486            AvroStreamOpts::builder()
487                .source(IngestSource::S3(S3IngestSource::new(s3_path)))
488                .target(DatasetIngestTarget::Existing(
489                    ExistingDatasetIngestDestination::new(data_source_rid),
490                ))
491                .build(),
492        );
493
494        let request = IngestRequest::new(opts);
495
496        self.ingest_client
497            .ingest(token, &request)
498            .await
499            .map_err(|e| UploaderError::Conjure(format!("{e:?}")))
500    }
501}
502
503pub struct ChunkedStreamReader {
504    reader: Box<dyn Read + Send>,
505    chunk_size: usize,
506    current_index: usize,
507}
508
509impl ChunkedStreamReader {
510    pub fn new<R>(reader: R, chunk_size: usize) -> Self
511    where
512        R: Read + Send + 'static,
513    {
514        Self {
515            reader: Box::new(reader),
516            chunk_size,
517            current_index: 0,
518        }
519    }
520}
521
522impl Stream for ChunkedStreamReader {
523    type Item = Result<(usize, Vec<u8>), std::io::Error>;
524
525    fn poll_next(
526        mut self: Pin<&mut Self>,
527        _cx: &mut std::task::Context<'_>,
528    ) -> std::task::Poll<Option<Self::Item>> {
529        let mut buffer = vec![0u8; self.chunk_size];
530
531        match self.reader.read(&mut buffer) {
532            Ok(0) => std::task::Poll::Ready(None),
533            Ok(n) => {
534                buffer.truncate(n);
535                let index = self.current_index;
536                self.current_index += 1;
537                std::task::Poll::Ready(Some(Ok((index, buffer))))
538            }
539            Err(e) => std::task::Poll::Ready(Some(Err(e))),
540        }
541    }
542}