1pub mod local;
13
14pub mod s3;
20
21#[cfg(test)]
22mod tests;
23
24use std::{
25 iter::{from_fn, once},
26 ops::Range,
27 sync::Arc,
28 time::Duration,
29};
30
31use bytes::Bytes;
32use derive_more::Debug;
33use futures::{FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt};
34use http::Method;
35use object_store::{
36 Attributes, CopyMode, DynObjectStore, GetResult, MultipartUpload, ObjectMeta, ObjectStore,
37 ObjectStoreExt, PutPayload, PutResult, path::Path, signer::Signer,
38};
39use tuwunel_core::{
40 Error, Result,
41 config::StorageProvider,
42 debug, err, error,
43 error::error_chain,
44 extract_variant, implement, info, trace,
45 utils::{
46 BoolExt,
47 result::FlatOk,
48 stream::{IterStream, TryReadyExt},
49 },
50};
51use url::Url;
52
53#[derive(Debug)]
59pub struct Provider {
60 pub name: String,
62
63 pub config: StorageProvider,
65
66 pub(crate) provider: Box<DynObjectStore>,
68
69 #[debug(skip)]
70 pub(crate) signer: Option<Arc<dyn Signer>>,
72
73 pub(crate) base_path: Option<Path>,
75
76 startup_check: bool,
77
78 #[expect(unused)]
79 #[debug(skip)]
80 services: Arc<crate::services::OnceServices>,
81}
82
83pub type FetchItem = (Bytes, (Range<u64>, u64));
89
90pub type FetchMetaItem = (Bytes, Arc<(Range<u64>, ObjectMeta, Attributes)>);
96
97#[implement(Provider)]
102#[tracing::instrument(skip_all, err)]
103pub(super) async fn start(self: &Arc<Self>) -> Result {
104 if self.startup_check {
105 self.startup_check().await?;
106 }
107
108 Ok(())
109}
110
111#[implement(Provider)]
112#[tracing::instrument(name = "check", skip_all, err)]
113async fn startup_check(self: &Arc<Self>) -> Result {
114 debug!(
115 name = ?self.name,
116 "Checking storage provider client connection...",
117 );
118 self.ping()
119 .inspect_ok(|()| {
120 info!(
121 name = %self.name,
122 "Connected to storage provider"
123 );
124 })
125 .await
126}
127
128#[implement(Provider)]
133#[tracing::instrument(
134 level = "debug",
135 err(level = "debug"),
136 skip_all,
137 fields(
138 provider = %self.name,
139 ?path,
140 ?size,
141 )
142)]
143pub async fn put<S, T>(&self, path: &str, size: Option<usize>, input: S) -> Result<PutResult>
144where
145 S: Stream<Item = Result<T>> + Send,
146 PutPayload: From<T> + From<PutPayload>,
147{
148 if size.is_none_or(|size| size >= self.multipart_threshold()) {
149 return self.put_multi(path, input).await;
150 }
151
152 debug!(
153 ?size,
154 threshold = ?self.multipart_threshold(),
155 "Selecting single-part upload..."
156 );
157
158 let payload: PutPayload = input
159 .map_ok(PutPayload::from)
160 .try_collect::<Vec<_>>()
161 .await?
162 .into_iter()
163 .map(Bytes::from)
164 .collect();
165
166 self.put_single(path, payload).await
167}
168
169#[implement(Provider)]
174#[tracing::instrument(
175 level = "debug",
176 err(level = "debug"),
177 skip_all,
178 fields(
179 provider = %self.name,
180 ?path,
181 )
182)]
183pub async fn put_one<T>(&self, path: &str, input: T) -> Result<PutResult>
184where
185 PutPayload: From<T> + From<PutPayload>,
186{
187 let payload: PutPayload = input.into();
188
189 if payload.content_length() < self.multipart_threshold() {
190 return self.put_single(path, payload).await;
191 }
192
193 let part_size = self.multipart_part_size();
194
195 debug!(
196 len = ?payload.content_length(),
197 threshold = ?self.multipart_threshold(),
198 ?part_size,
199 "Selecting multi-part upload..."
200 );
201
202 self.put_multi(path, chunked(payload, part_size).try_stream())
203 .await
204}
205
206#[implement(Provider)]
211#[tracing::instrument(
212 level = "debug",
213 err(level = "debug"),
214 skip_all,
215 fields(
216 provider = %self.name,
217 ?path,
218 )
219)]
220async fn put_multi<S, T>(&self, path: &str, input: S) -> Result<PutResult>
221where
222 S: Stream<Item = Result<T>> + Send,
223 PutPayload: From<T> + From<PutPayload>,
224{
225 let path = self.to_abs_path(path)?;
226 let mut handle = self
227 .provider
228 .put_multipart(&path)
229 .map_err(Error::from)
230 .await?;
231
232 match input
233 .try_for_each(|t| handle.put_part(t.into()).map_err(Error::from))
234 .inspect_err(|e| error!(?path, chain = %error_chain(e), "Failed to store object"))
235 .await
236 {
237 | Ok(()) =>
238 handle
239 .complete()
240 .map_err(Error::from)
241 .inspect_err(|e| {
242 error!(
243 ?path,
244 chain = %error_chain(e),
245 "Failed to store object during completion",
246 );
247 })
248 .await,
249
250 | Err(e) =>
251 handle
252 .abort()
253 .map_ok(|()| Err(e))
254 .map_err(Error::from)
255 .inspect_err(|e| {
256 error!(
257 ?path,
258 chain = %error_chain(e),
259 "Additional errors during error handling",
260 );
261 })
262 .await?,
263 }
264}
265
266#[implement(Provider)]
271#[tracing::instrument(
272 level = "debug",
273 err(level = "debug"),
274 skip_all,
275 fields(
276 provider = %self.name,
277 ?path,
278 )
279)]
280async fn put_single(&self, path: &str, input: PutPayload) -> Result<PutResult> {
281 let path = self.to_abs_path(path)?;
282
283 self.provider
284 .put(&path, input)
285 .map_err(Error::from)
286 .await
287}
288
289#[implement(Provider)]
295#[tracing::instrument(
296 level = "debug",
297 skip_all,
298 fields(
299 provider = %self.name,
300 ?path,
301 )
302)]
303pub fn fetch_with_metadata(
304 &self,
305 path: &str,
306) -> impl Stream<Item = Result<FetchMetaItem>> + Send {
307 self.load(path)
308 .map_ok(|result| {
309 let meta = (result.range.clone(), result.meta.clone(), result.attributes.clone());
310 let data = Arc::new(meta);
311
312 result
313 .into_stream()
314 .map_err(Error::from)
315 .map_ok(move |bytes| (bytes, data.clone()))
316 })
317 .map_err(Error::from)
318 .try_flatten_stream()
319}
320
321#[implement(Provider)]
327#[tracing::instrument(
328 level = "debug",
329 skip_all,
330 fields(
331 provider = %self.name,
332 ?path,
333 )
334)]
335pub fn fetch(&self, path: &str) -> impl Stream<Item = Result<FetchItem>> + Send {
336 self.load(path)
337 .map_ok(|result| {
338 let size = result.meta.size;
339 let range = result.range.clone();
340
341 result
342 .into_stream()
343 .map_err(Error::from)
344 .map_ok(move |bytes| (bytes, (range.clone(), size)))
345 })
346 .map_err(Error::from)
347 .try_flatten_stream()
348}
349
350#[implement(Provider)]
355#[tracing::instrument(
356 level = "debug",
357 err(level = "debug"),
358 skip_all,
359 fields(
360 provider = %self.name,
361 ?path,
362 )
363)]
364pub async fn get(&self, path: &str) -> Result<Bytes> {
365 self.load(path)
366 .map_ok(GetResult::bytes)
367 .await?
368 .map_err(Error::from)
369 .await
370}
371
372#[implement(Provider)]
377#[tracing::instrument(
378 level = "debug",
379 err(level = "debug"),
380 skip_all,
381 fields(
382 provider = %self.name,
383 ?path,
384 )
385)]
386pub async fn load(&self, path: &str) -> Result<GetResult> {
387 let path = self.to_abs_path(path)?;
388
389 self.provider
390 .get(&path)
391 .map_err(Error::from)
392 .await
393}
394
395#[implement(Provider)]
400#[tracing::instrument(
401 level = "debug",
402 err(level = "debug"),
403 skip_all,
404 fields(
405 provider = %self.name,
406 ?path,
407 ?ttl,
408 )
409)]
410pub async fn signed_get_url(&self, path: &str, ttl: Duration) -> Result<Option<Url>> {
411 let Some(signer) = self.signer.as_ref() else {
412 return Ok(None);
413 };
414
415 let path = self.to_abs_path(path)?;
416
417 signer
418 .signed_url(Method::GET, &path, ttl)
419 .map_err(Error::from)
420 .map_ok(Some)
421 .await
422}
423
424#[implement(Provider)]
429#[tracing::instrument(
430 level = "debug",
431 err(level = "debug"),
432 skip_all,
433 fields(
434 provider = %self.name,
435 ?path,
436 )
437)]
438pub async fn delete_one(self: &Arc<Self>, path: &str) -> Result {
439 self.delete(once(path.to_owned()).stream())
440 .map_ok(|_| ())
441 .try_collect()
442 .await
443}
444
445#[implement(Provider)]
450#[tracing::instrument(
451 level = "debug",
452 skip_all,
453 fields(
454 provider = %self.name,
455 )
456)]
457pub fn delete<S>(self: &Arc<Self>, paths: S) -> impl Stream<Item = Result<Path>> + Send
458where
459 S: Stream<Item = String> + Send + 'static,
460{
461 let this = self.clone();
462 let paths = paths
463 .map(Ok)
464 .ready_and_then(move |path| {
465 use object_store::{Error, path};
466
467 this.to_abs_path(&path)
468 .map_err(|_| Error::InvalidPath {
469 source: path::Error::InvalidPath { path: path.into() },
470 })
471 })
472 .boxed();
473
474 self.provider
475 .delete_stream(paths)
476 .map_err(Error::from)
477}
478
479#[implement(Provider)]
484#[tracing::instrument(
485 level = "debug",
486 err(level = "debug"),
487 skip_all,
488 fields(
489 provider = %self.name,
490 ?src,
491 ?dst,
492 ?overwrite,
493 )
494)]
495pub async fn rename(&self, src: &str, dst: &str, overwrite: CopyMode) -> Result {
496 let src = self.to_abs_path(src)?;
497 let dst = self.to_abs_path(dst)?;
498
499 match overwrite {
500 | CopyMode::Overwrite => self.provider.rename(&src, &dst).left_future(),
501 | CopyMode::Create => self
502 .provider
503 .rename_if_not_exists(&src, &dst)
504 .right_future(),
505 }
506 .map_err(Error::from)
507 .await
508}
509
510#[implement(Provider)]
515#[tracing::instrument(
516 level = "debug",
517 err(level = "debug"),
518 skip_all,
519 fields(
520 provider = %self.name,
521 ?src,
522 ?dst,
523 ?overwrite,
524 )
525)]
526pub async fn copy(&self, src: &str, dst: &str, overwrite: CopyMode) -> Result {
527 let src = self.to_abs_path(src)?;
528 let dst = self.to_abs_path(dst)?;
529
530 match overwrite {
531 | CopyMode::Overwrite => self.provider.copy(&src, &dst).left_future(),
532 | CopyMode::Create => self
533 .provider
534 .copy_if_not_exists(&src, &dst)
535 .right_future(),
536 }
537 .map_err(Error::from)
538 .await
539}
540
541#[implement(Provider)]
546#[tracing::instrument(
547 level = "debug",
548 skip_all,
549 fields(
550 provider = %self.name,
551 ?prefix,
552 )
553)]
554pub fn list(&self, prefix: Option<&str>) -> impl Stream<Item = Result<ObjectMeta>> + Send {
555 let abs_prefix = prefix
556 .map(Path::from)
557 .map(|p| self.prepend_base_path(p))
558 .or_else(|| self.base_path.clone());
559
560 self.provider
561 .list(abs_prefix.as_ref())
562 .map_err(Error::from)
563 .map_ok(|meta| ObjectMeta {
564 location: self.strip_base_path(meta.location),
565 ..meta
566 })
567}
568
569#[implement(Provider)]
574#[tracing::instrument(
575 level = "debug",
576 err(level = "debug"),
577 skip_all,
578 fields(
579 provider = %self.name,
580 ?path,
581 )
582)]
583pub async fn head(&self, path: &str) -> Result<ObjectMeta> {
584 self.provider
585 .head(&self.to_abs_path(path)?)
586 .map_err(Error::from)
587 .await
588}
589
590#[implement(Provider)]
595#[tracing::instrument(
596 level = "debug",
597 err(level = "error"),
598 skip_all,
599 fields(
600 provider = %self.name,
601 )
602)]
603pub async fn ping(&self) -> Result {
604 self.list(None)
605 .try_next()
606 .inspect_err(|e| {
607 error!(chain = %error_chain(e), "Failed to connect to storage provider");
608 })
609 .boxed()
610 .await
611 .map(|_| ())
612}
613
614#[implement(Provider)]
615fn to_abs_path(&self, location: &str) -> Result<Path> {
616 let location = Path::parse(location)
617 .map_err(|e| err!("Failed to parse location into canonical PathPart: {e}"))?;
618
619 let path = self.prepend_base_path(location);
620
621 trace!(
622 provider = ?self.name,
623 base_path = ?self.base_path,
624 ?path,
625 "Computed absolute path for object on provider.",
626 );
627
628 Ok(path)
629}
630
631#[implement(Provider)]
632fn prepend_base_path(&self, location: Path) -> Path {
633 match self.base_path.as_ref() {
634 | Some(base_path) if !location.prefix_matches(base_path) => base_path
635 .parts()
636 .chain(location.parts())
637 .collect(),
638
639 | _ => location,
640 }
641}
642
643#[implement(Provider)]
644fn strip_base_path(&self, location: Path) -> Path {
645 self.base_path
646 .as_ref()
647 .and_then(|base_path| location.prefix_match(base_path))
648 .map(Iterator::collect)
649 .unwrap_or(location)
650}
651
652#[implement(Provider)]
653fn multipart_threshold(&self) -> usize {
654 extract_variant!(&self.config, StorageProvider::s3)
655 .map(|config| config.multipart_threshold.as_u64())
656 .map(TryInto::try_into)
657 .flat_ok()
658 .unwrap_or(usize::MAX)
659}
660
661#[implement(Provider)]
662fn multipart_part_size(&self) -> usize {
663 extract_variant!(&self.config, StorageProvider::s3)
664 .map(|config| config.multipart_part_size.as_u64())
665 .map(TryInto::try_into)
666 .flat_ok()
667 .unwrap_or(usize::MAX)
668}
669
670fn chunked(payload: PutPayload, part_size: usize) -> impl Iterator<Item = PutPayload> {
675 let mut buf: Bytes = payload.into();
676 from_fn(move || {
677 buf.is_empty()
678 .is_false()
679 .then(|| buf.split_to(part_size.min(buf.len())).into())
680 })
681}