Skip to main content

shared/
storage.rs

1use crate::settings::SettingsReadGuard;
2use aws_sdk_s3::{
3    Client as S3Client,
4    config::{Config as S3Config, Credentials, Region, retry::RetryConfig, timeout::TimeoutConfig},
5    primitives::ByteStream,
6    types::{CompletedMultipartUpload, CompletedPart},
7};
8use compact_str::ToCompactString;
9use serde::{Deserialize, Serialize};
10use std::{path::Path, sync::Arc};
11use tokio::io::{AsyncReadExt, AsyncWriteExt};
12use tokio_util::bytes::{Bytes, BytesMut};
13use utoipa::ToSchema;
14
15#[derive(ToSchema, Deserialize, Serialize)]
16pub struct StorageAsset {
17    pub name: compact_str::CompactString,
18    pub url: String,
19    pub size: u64,
20    pub is_directory: bool,
21    pub created: chrono::DateTime<chrono::Utc>,
22}
23
24const SEARCH_SCAN_LIMIT: usize = 10_000;
25
26fn get_s3_client(
27    access_key: &str,
28    secret_key: &str,
29    region: &str,
30    endpoint: &str,
31    path_style: bool,
32) -> Result<S3Client, anyhow::Error> {
33    let credentials = Credentials::new(access_key, secret_key, None, None, "calagopus-static");
34
35    let timeout_config = TimeoutConfig::builder()
36        .connect_timeout(std::time::Duration::from_secs(10))
37        .build();
38
39    let config = S3Config::builder()
40        .behavior_version(aws_sdk_s3::config::BehaviorVersion::latest())
41        .credentials_provider(credentials)
42        .region(Region::new(region.to_string()))
43        .endpoint_url(endpoint)
44        .force_path_style(path_style)
45        .timeout_config(timeout_config)
46        .retry_config(RetryConfig::standard())
47        .build();
48
49    Ok(S3Client::from_conf(config))
50}
51
52pub struct StorageUrlRetriever<'a> {
53    settings: SettingsReadGuard<'a>,
54}
55
56impl<'a> StorageUrlRetriever<'a> {
57    pub fn new(settings: SettingsReadGuard<'a>) -> Self {
58        Self { settings }
59    }
60
61    pub fn get_settings(&self) -> &super::settings::AppSettings {
62        &self.settings
63    }
64
65    pub fn get_url(&self, path: impl AsRef<str>) -> String {
66        match &self.settings.storage_driver {
67            super::settings::StorageDriver::Filesystem { .. } => {
68                format!(
69                    "{}/{}",
70                    self.settings.app.url.trim_end_matches('/'),
71                    path.as_ref()
72                )
73            }
74            super::settings::StorageDriver::S3 { public_url, .. } => {
75                format!("{}/{}", public_url.trim_end_matches('/'), path.as_ref())
76            }
77        }
78    }
79}
80
81pub struct Storage {
82    settings: Arc<super::settings::Settings>,
83}
84
85impl Storage {
86    pub fn new(settings: Arc<super::settings::Settings>) -> Self {
87        Self { settings }
88    }
89
90    pub async fn retrieve_urls(&self) -> Result<StorageUrlRetriever<'_>, anyhow::Error> {
91        let settings = self.settings.get().await?;
92
93        Ok(StorageUrlRetriever::new(settings))
94    }
95
96    pub async fn remove(&self, path: Option<impl AsRef<str>>) -> Result<(), anyhow::Error> {
97        let path = match path {
98            Some(path) => path,
99            None => return Ok(()),
100        };
101        let path = path.as_ref();
102
103        if path.is_empty() || path.contains("..") || path.starts_with('/') {
104            return Err(anyhow::anyhow!("invalid path"));
105        }
106
107        let settings = self.settings.get().await?;
108
109        tracing::debug!(path, "removing file");
110
111        match &settings.storage_driver {
112            super::settings::StorageDriver::Filesystem { path: base_path } => {
113                let base_filesystem =
114                    match crate::cap::CapFilesystem::async_new(base_path.into()).await {
115                        Ok(base_filesystem) => base_filesystem,
116                        Err(err) if err.kind() == std::io::ErrorKind::NotFound => return Ok(()),
117                        Err(err) => return Err(err.into()),
118                    };
119                drop(settings);
120
121                if let Err(err) = base_filesystem.async_remove_file(&path).await
122                    && err
123                        .downcast_ref::<std::io::Error>()
124                        .is_none_or(|e| e.kind() != std::io::ErrorKind::NotFound)
125                {
126                    return Err(err);
127                }
128
129                if let Some(parent) = Path::new(path).parent().map(|p| p.to_path_buf()) {
130                    tokio::spawn(async move {
131                        tokio::time::sleep(std::time::Duration::from_secs(10)).await;
132
133                        let mut directory = match base_filesystem.async_read_dir(&parent).await {
134                            Ok(directory) => directory,
135                            Err(_) => return,
136                        };
137
138                        if directory.next_entry().await.is_none() {
139                            base_filesystem.async_remove_dir(parent).await.ok();
140                        }
141                    });
142                }
143            }
144            super::settings::StorageDriver::S3 {
145                access_key,
146                secret_key,
147                bucket,
148                region,
149                endpoint,
150                path_style,
151                ..
152            } => {
153                let s3_client =
154                    get_s3_client(access_key, secret_key, region, endpoint, *path_style)?;
155                let bucket = bucket.clone();
156                drop(settings);
157
158                s3_client
159                    .delete_object()
160                    .bucket(bucket)
161                    .key(path)
162                    .send()
163                    .await?;
164            }
165        }
166
167        Ok(())
168    }
169
170    pub async fn store(
171        &self,
172        path: impl AsRef<str>,
173        data: impl tokio::io::AsyncRead + Unpin,
174        content_type: impl AsRef<str>,
175    ) -> Result<u64, anyhow::Error> {
176        let path = path.as_ref();
177        let content_type = content_type.as_ref();
178
179        if path.is_empty() || path.contains("..") || path.starts_with('/') {
180            return Err(anyhow::anyhow!("invalid path"));
181        }
182
183        let settings = self.settings.get().await?;
184
185        tracing::debug!(path, content_type, "storing file");
186
187        match &settings.storage_driver {
188            super::settings::StorageDriver::Filesystem { path: base_path } => {
189                tokio::fs::create_dir_all(base_path).await?;
190
191                let base_filesystem =
192                    crate::cap::CapFilesystem::async_new(base_path.into()).await?;
193                drop(settings);
194
195                if let Some(parent) = Path::new(path).parent() {
196                    base_filesystem.async_create_dir_all(parent).await?;
197                }
198
199                let mut file = base_filesystem.async_create(path).await?;
200                let mut data = data;
201                let bytes = tokio::io::copy(&mut data, &mut file).await?;
202
203                file.shutdown().await?;
204                Ok(bytes)
205            }
206            super::settings::StorageDriver::S3 {
207                access_key,
208                secret_key,
209                bucket,
210                region,
211                endpoint,
212                path_style,
213                ..
214            } => {
215                let s3_client =
216                    get_s3_client(access_key, secret_key, region, endpoint, *path_style)?;
217                let bucket = bucket.clone();
218                drop(settings);
219
220                upload_multipart(&s3_client, &bucket, path, content_type, data).await
221            }
222        }
223    }
224
225    pub async fn list(
226        &self,
227        base: impl AsRef<str>,
228        directory: impl AsRef<str>,
229        page: usize,
230        per_page: usize,
231    ) -> Result<crate::models::Pagination<StorageAsset>, anyhow::Error> {
232        let base = base.as_ref();
233        let directory = directory.as_ref();
234
235        if base.is_empty() || base.contains("..") || base.starts_with('/') {
236            return Err(anyhow::anyhow!("invalid base path"));
237        }
238        if !directory.is_empty()
239            && (directory.contains("..") || directory.starts_with('/') || directory.ends_with('/'))
240        {
241            return Err(anyhow::anyhow!("invalid directory path"));
242        }
243
244        let settings = self.settings.get().await?;
245
246        match &settings.storage_driver {
247            super::settings::StorageDriver::Filesystem { path: base_path } => {
248                let dir_path = if directory.is_empty() {
249                    Path::new(base_path).join(base)
250                } else {
251                    Path::new(base_path).join(base).join(directory)
252                };
253
254                let base_filesystem = match crate::cap::CapFilesystem::async_new(dir_path).await {
255                    Ok(base_filesystem) => base_filesystem,
256                    Err(err) if err.kind() == std::io::ErrorKind::NotFound => {
257                        return Ok(crate::models::Pagination {
258                            total: 0,
259                            per_page: per_page as i64,
260                            page: page as i64,
261                            data: Vec::new(),
262                        });
263                    }
264                    Err(err) => return Err(err.into()),
265                };
266                drop(settings);
267
268                let mut dir_reader = base_filesystem.async_read_dir("").await?;
269                let mut raw_dirs: Vec<String> = Vec::new();
270                let mut raw_files: Vec<String> = Vec::new();
271
272                while let Some(Ok((is_dir, name))) = dir_reader.next_entry().await {
273                    if is_dir {
274                        raw_dirs.push(name);
275                    } else {
276                        raw_files.push(name);
277                    }
278                }
279
280                raw_dirs.sort_unstable();
281                raw_files.sort_unstable();
282
283                let total = (raw_dirs.len() + raw_files.len()) as i64;
284                let start = (page - 1) * per_page;
285
286                let storage_url_retriever = self.retrieve_urls().await?;
287
288                let mut entries = Vec::new();
289
290                for (is_dir, name) in raw_dirs
291                    .into_iter()
292                    .map(|n| (true, n))
293                    .chain(raw_files.into_iter().map(|n| (false, n)))
294                    .skip(start)
295                    .take(per_page)
296                {
297                    let full_name = if directory.is_empty() {
298                        name.clone()
299                    } else {
300                        format!("{directory}/{name}")
301                    };
302
303                    let (size, created) = if is_dir {
304                        (0u64, chrono::DateTime::<chrono::Utc>::default())
305                    } else {
306                        let metadata = match base_filesystem.async_metadata(&name).await {
307                            Ok(m) => m,
308                            Err(_) => continue,
309                        };
310                        let created = metadata
311                            .created()
312                            .or_else(|_| metadata.modified())?
313                            .into_std()
314                            .into();
315                        (metadata.len(), created)
316                    };
317
318                    entries.push(StorageAsset {
319                        url: storage_url_retriever.get_url(format!("{base}/{full_name}")),
320                        name: full_name.to_compact_string(),
321                        size,
322                        is_directory: is_dir,
323                        created,
324                    });
325                }
326
327                Ok(crate::models::Pagination {
328                    total,
329                    per_page: per_page as i64,
330                    page: page as i64,
331                    data: entries,
332                })
333            }
334            super::settings::StorageDriver::S3 {
335                access_key,
336                secret_key,
337                bucket,
338                region,
339                endpoint,
340                path_style,
341                ..
342            } => {
343                let s3_client =
344                    get_s3_client(access_key, secret_key, region, endpoint, *path_style)?;
345                let bucket_name = bucket.clone();
346                drop(settings);
347
348                let s3_prefix = if directory.is_empty() {
349                    format!("{base}/")
350                } else {
351                    format!("{base}/{directory}/")
352                };
353                let strip_prefix = format!("{base}/");
354
355                let storage_url_retriever = self.retrieve_urls().await?;
356
357                let mut dirs = Vec::new();
358                let mut files = Vec::new();
359
360                let mut paginator = s3_client
361                    .list_objects_v2()
362                    .bucket(&*bucket_name)
363                    .prefix(&s3_prefix)
364                    .delimiter("/")
365                    .into_paginator()
366                    .send();
367
368                while let Some(result) = paginator.next().await {
369                    let page = result?;
370
371                    for cp in page.common_prefixes() {
372                        let Some(prefix) = cp.prefix() else { continue };
373                        let name = prefix
374                            .strip_prefix(&strip_prefix)
375                            .unwrap_or(prefix)
376                            .trim_end_matches('/')
377                            .to_compact_string();
378                        dirs.push(StorageAsset {
379                            url: storage_url_retriever.get_url(prefix),
380                            name,
381                            size: 0,
382                            is_directory: true,
383                            created: chrono::DateTime::<chrono::Utc>::default(),
384                        });
385                    }
386
387                    for entry in page.contents() {
388                        let Some(key) = entry.key() else { continue };
389                        if key == s3_prefix {
390                            continue;
391                        }
392                        let name = key
393                            .strip_prefix(&strip_prefix)
394                            .unwrap_or(key)
395                            .to_compact_string();
396                        let size = entry.size().unwrap_or(0).max(0) as u64;
397                        let created = entry
398                            .last_modified()
399                            .and_then(|dt| {
400                                chrono::DateTime::<chrono::Utc>::from_timestamp(
401                                    dt.secs(),
402                                    dt.subsec_nanos(),
403                                )
404                            })
405                            .unwrap_or_default();
406
407                        files.push(StorageAsset {
408                            url: storage_url_retriever.get_url(key),
409                            name,
410                            size,
411                            is_directory: false,
412                            created,
413                        });
414                    }
415                }
416
417                let total = (dirs.len() + files.len()) as i64;
418                let start = (page - 1) * per_page;
419
420                Ok(crate::models::Pagination {
421                    total,
422                    per_page: per_page as i64,
423                    page: page as i64,
424                    data: dirs
425                        .into_iter()
426                        .chain(files)
427                        .skip(start)
428                        .take(per_page)
429                        .collect(),
430                })
431            }
432        }
433    }
434
435    pub async fn search(
436        &self,
437        base: impl AsRef<str>,
438        directory: impl AsRef<str>,
439        search: impl AsRef<str>,
440        limit: usize,
441    ) -> Result<Vec<StorageAsset>, anyhow::Error> {
442        let base = base.as_ref();
443        let directory = directory.as_ref();
444        let search = search.as_ref().trim().to_lowercase();
445
446        if base.is_empty() || base.contains("..") || base.starts_with('/') {
447            return Err(anyhow::anyhow!("invalid base path"));
448        }
449        if !directory.is_empty()
450            && (directory.contains("..") || directory.starts_with('/') || directory.ends_with('/'))
451        {
452            return Err(anyhow::anyhow!("invalid directory path"));
453        }
454        if search.is_empty() {
455            return Err(anyhow::anyhow!("empty search"));
456        }
457
458        let settings = self.settings.get().await?;
459
460        let mut dirs = Vec::new();
461        let mut files = Vec::new();
462        let mut scanned = 0;
463
464        match &settings.storage_driver {
465            super::settings::StorageDriver::Filesystem { path: base_path } => {
466                let dir_path = if directory.is_empty() {
467                    Path::new(base_path).join(base)
468                } else {
469                    Path::new(base_path).join(base).join(directory)
470                };
471
472                let base_filesystem = match crate::cap::CapFilesystem::async_new(dir_path).await {
473                    Ok(base_filesystem) => base_filesystem,
474                    Err(err) if err.kind() == std::io::ErrorKind::NotFound => {
475                        return Ok(Vec::new());
476                    }
477                    Err(err) => return Err(err.into()),
478                };
479                drop(settings);
480
481                let storage_url_retriever = self.retrieve_urls().await?;
482                let mut walker = base_filesystem.async_walk_dir("").await?;
483
484                while let Some(Ok((is_dir, path))) = walker.next_entry().await {
485                    scanned += 1;
486                    if scanned > SEARCH_SCAN_LIMIT {
487                        break;
488                    }
489
490                    let Some(name) = path.file_name().and_then(|name| name.to_str()) else {
491                        continue;
492                    };
493                    if !name.to_lowercase().contains(&search) {
494                        continue;
495                    }
496
497                    let relative = path.to_string_lossy();
498                    let full_name = if directory.is_empty() {
499                        relative.to_compact_string()
500                    } else {
501                        format!("{directory}/{relative}").to_compact_string()
502                    };
503
504                    let (size, created) = if is_dir {
505                        (0u64, chrono::DateTime::<chrono::Utc>::default())
506                    } else {
507                        let metadata = match base_filesystem.async_metadata(&path).await {
508                            Ok(m) => m,
509                            Err(_) => continue,
510                        };
511                        let created = metadata
512                            .created()
513                            .or_else(|_| metadata.modified())?
514                            .into_std()
515                            .into();
516                        (metadata.len(), created)
517                    };
518
519                    let asset = StorageAsset {
520                        url: storage_url_retriever.get_url(format!("{base}/{full_name}")),
521                        name: full_name,
522                        size,
523                        is_directory: is_dir,
524                        created,
525                    };
526
527                    if is_dir { &mut dirs } else { &mut files }.push(asset);
528                }
529            }
530            super::settings::StorageDriver::S3 {
531                access_key,
532                secret_key,
533                bucket,
534                region,
535                endpoint,
536                path_style,
537                ..
538            } => {
539                let s3_client =
540                    get_s3_client(access_key, secret_key, region, endpoint, *path_style)?;
541                let bucket_name = bucket.clone();
542                drop(settings);
543
544                let s3_prefix = if directory.is_empty() {
545                    format!("{base}/")
546                } else {
547                    format!("{base}/{directory}/")
548                };
549                let strip_prefix = format!("{base}/");
550
551                let storage_url_retriever = self.retrieve_urls().await?;
552                let mut seen_dirs = std::collections::HashSet::new();
553
554                let mut paginator = s3_client
555                    .list_objects_v2()
556                    .bucket(&*bucket_name)
557                    .prefix(&s3_prefix)
558                    .into_paginator()
559                    .send();
560
561                'scan: while let Some(result) = paginator.next().await {
562                    let objects = result?;
563
564                    for entry in objects.contents() {
565                        let Some(key) = entry.key() else { continue };
566                        if key.ends_with('/') {
567                            continue;
568                        }
569
570                        scanned += 1;
571                        if scanned > SEARCH_SCAN_LIMIT {
572                            break 'scan;
573                        }
574
575                        let relative = key.strip_prefix(&strip_prefix).unwrap_or(key);
576
577                        let mut segment_start = 0;
578                        for (index, char) in relative.char_indices() {
579                            if char != '/' {
580                                continue;
581                            }
582
583                            let dir_path = &relative[..index];
584                            let name = &relative[segment_start..index];
585                            segment_start = index + 1;
586
587                            if dir_path.len() <= directory.len()
588                                || !name.to_lowercase().contains(&search)
589                                || !seen_dirs.insert(dir_path.to_compact_string())
590                            {
591                                continue;
592                            }
593
594                            dirs.push(StorageAsset {
595                                url: storage_url_retriever.get_url(format!("{base}/{dir_path}/")),
596                                name: dir_path.to_compact_string(),
597                                size: 0,
598                                is_directory: true,
599                                created: chrono::DateTime::<chrono::Utc>::default(),
600                            });
601                        }
602
603                        if !relative[segment_start..].to_lowercase().contains(&search) {
604                            continue;
605                        }
606
607                        files.push(StorageAsset {
608                            url: storage_url_retriever.get_url(key),
609                            name: relative.to_compact_string(),
610                            size: entry.size().unwrap_or(0).max(0) as u64,
611                            created: entry
612                                .last_modified()
613                                .and_then(|dt| {
614                                    chrono::DateTime::<chrono::Utc>::from_timestamp(
615                                        dt.secs(),
616                                        dt.subsec_nanos(),
617                                    )
618                                })
619                                .unwrap_or_default(),
620                            is_directory: false,
621                        });
622                    }
623                }
624            }
625        }
626
627        dirs.sort_unstable_by(|a: &StorageAsset, b: &StorageAsset| a.name.cmp(&b.name));
628        files.sort_unstable_by(|a: &StorageAsset, b: &StorageAsset| a.name.cmp(&b.name));
629
630        Ok(dirs.into_iter().chain(files).take(limit).collect())
631    }
632}
633
634const PART_SIZE: usize = 16 * 1024 * 1024;
635
636async fn upload_multipart(
637    client: &S3Client,
638    bucket: &str,
639    key: &str,
640    content_type: &str,
641    mut data: impl tokio::io::AsyncRead + Unpin,
642) -> Result<u64, anyhow::Error> {
643    let first_part = read_part(&mut data, PART_SIZE).await?;
644
645    if first_part.len() < PART_SIZE {
646        let total = first_part.len() as u64;
647        client
648            .put_object()
649            .bucket(bucket)
650            .key(key)
651            .content_type(content_type)
652            .body(ByteStream::from(first_part))
653            .send()
654            .await?;
655        return Ok(total);
656    }
657
658    let create = client
659        .create_multipart_upload()
660        .bucket(bucket)
661        .key(key)
662        .content_type(content_type)
663        .send()
664        .await?;
665
666    let upload_id = create
667        .upload_id()
668        .ok_or_else(|| anyhow::anyhow!("S3 did not return an upload_id"))?
669        .to_string();
670
671    let result = run_multipart(client, bucket, key, &upload_id, &mut data, first_part).await;
672
673    match result {
674        Ok(total) => Ok(total),
675        Err(err) => {
676            if let Err(abort_err) = client
677                .abort_multipart_upload()
678                .bucket(bucket)
679                .key(key)
680                .upload_id(&upload_id)
681                .send()
682                .await
683            {
684                tracing::warn!(
685                    bucket,
686                    key,
687                    upload_id,
688                    "failed to abort multipart upload after error: {:#?}",
689                    abort_err
690                );
691            }
692            Err(err)
693        }
694    }
695}
696
697async fn run_multipart(
698    client: &S3Client,
699    bucket: &str,
700    key: &str,
701    upload_id: &str,
702    data: &mut (impl tokio::io::AsyncRead + Unpin),
703    first_part: Bytes,
704) -> Result<u64, anyhow::Error> {
705    let mut completed = Vec::new();
706    let mut total: u64 = 0;
707    let mut part_number: i32 = 1;
708    let mut current = first_part;
709
710    loop {
711        let part_len = current.len() as u64;
712        let resp = client
713            .upload_part()
714            .bucket(bucket)
715            .key(key)
716            .upload_id(upload_id)
717            .part_number(part_number)
718            .body(ByteStream::from(current))
719            .send()
720            .await?;
721
722        completed.push(
723            CompletedPart::builder()
724                .part_number(part_number)
725                .set_e_tag(resp.e_tag().map(|s| s.to_string()))
726                .build(),
727        );
728        total += part_len;
729        part_number += 1;
730
731        let next = read_part(data, PART_SIZE).await?;
732        if next.is_empty() {
733            break;
734        }
735        current = next;
736    }
737
738    let completed_upload = CompletedMultipartUpload::builder()
739        .set_parts(Some(completed))
740        .build();
741
742    client
743        .complete_multipart_upload()
744        .bucket(bucket)
745        .key(key)
746        .upload_id(upload_id)
747        .multipart_upload(completed_upload)
748        .send()
749        .await?;
750
751    Ok(total)
752}
753
754async fn read_part(
755    data: &mut (impl tokio::io::AsyncRead + Unpin),
756    cap: usize,
757) -> Result<Bytes, std::io::Error> {
758    let mut buf = BytesMut::with_capacity(cap);
759    while buf.len() < cap {
760        let n = data.read_buf(&mut buf).await?;
761        if n == 0 {
762            break;
763        }
764    }
765    Ok(buf.freeze())
766}