daggerml.contrib.s3

View source
  1from __future__ import annotations
  2
  3import fnmatch
  4import hashlib
  5import io
  6import json
  7import os
  8import tarfile
  9from dataclasses import dataclass
 10from pathlib import Path
 11from typing import Any, Iterable, Literal, cast
 12from urllib.parse import urlparse
 13
 14from daggerml import Node, Uri, get_default_dml
 15from daggerml.api import DmlRepoError
 16from daggerml.util import get_client
 17
 18
 19def is_s3_uri(value: str) -> bool:
 20    p = urlparse(value)
 21    return p.scheme == "s3" and bool(p.netloc) and bool(p.path and p.path != "/")
 22
 23
 24def _sha256_bytes(data: bytes) -> str:
 25    return hashlib.sha256(data).hexdigest()
 26
 27
 28def _flatten_names(*name_or_uris):
 29    if len(name_or_uris) == 1 and isinstance(name_or_uris[0], (list, tuple)):
 30        return list(name_or_uris[0])
 31    return list(name_or_uris)
 32
 33
 34def _validate_safe_extract_path(*, dest_path: Path, member_name: str) -> None:
 35    member_path = Path(member_name)
 36    if member_path.is_absolute():
 37        raise DmlRepoError(f"Refusing to extract absolute tar path: {member_name}")
 38    target_path = (dest_path / member_path).resolve()
 39    if os.path.commonpath([str(dest_path), str(target_path)]) != str(dest_path):
 40        raise DmlRepoError(f"Refusing to extract path outside destination: {member_name}")
 41
 42
 43@dataclass(frozen=True)
 44class S3Store:
 45    bucket: str | None = None
 46    prefix: str | None = None
 47    client: Any = None
 48
 49    def __post_init__(self):
 50        bucket = self.bucket
 51        prefix = self.prefix
 52        if bucket is None and prefix is None:
 53            # use `daggerml`'s default dml session
 54            remote_root = get_default_dml().config.show()["remote"]["root"]
 55            if not remote_root:
 56                raise DmlRepoError(
 57                    "S3Store requires configured remote.root (set DML_REMOTE_ROOT or pass bucket/prefix)"
 58                )
 59            p = urlparse(remote_root)
 60            if p.scheme != "s3" or not p.netloc:
 61                raise DmlRepoError("remote.root must be an s3:// URI")
 62            bucket = p.netloc
 63            base = p.path.lstrip("/").rstrip("/")
 64            prefix = f"{base}/data" if base else "data"
 65        if bucket is None:
 66            raise DmlRepoError("S3Store bucket not configured")
 67        if prefix is None:
 68            prefix = ""
 69        object.__setattr__(self, "bucket", bucket)
 70        object.__setattr__(self, "prefix", prefix.strip("/"))
 71        object.__setattr__(self, "client", self.client or get_client("s3"))
 72
 73    @classmethod
 74    def from_remote_root(cls, remote_root: str) -> "S3Store":
 75        p = urlparse(remote_root)
 76        if p.scheme != "s3" or not p.netloc:
 77            raise DmlRepoError("remote root must be an s3:// URI")
 78        base = p.path.lstrip("/").rstrip("/")
 79        prefix = f"{base}/data" if base else "data"
 80        return cls(bucket=p.netloc, prefix=prefix)
 81
 82    def parse_uri(self, name_or_uri) -> tuple[str, str]:
 83        if isinstance(name_or_uri, Node):
 84            name_or_uri = name_or_uri.value()
 85        if isinstance(name_or_uri, Uri):
 86            name_or_uri = name_or_uri.uri
 87        if not isinstance(name_or_uri, str):
 88            raise DmlRepoError("S3Store name_or_uri must be a string or uri-bearing object")
 89        p = urlparse(name_or_uri)
 90        if p.scheme == "s3":
 91            return p.netloc, p.path[1:]
 92        if self.bucket is None:
 93            raise DmlRepoError("S3Store bucket not configured")
 94        key = f"{self.prefix}/{name_or_uri}" if self.prefix else name_or_uri
 95        return cast(str, self.bucket), key
 96
 97    def _name2uri(self, name) -> Uri:
 98        bucket, key = self.parse_uri(name)
 99        return Uri(f"s3://{bucket}/{key}")
100
101    def put(self, data: bytes | None = None, filepath: str | None = None, *, suffix: str = "") -> Uri:
102        if (data is None) == (filepath is None):
103            raise DmlRepoError("S3Store.put requires exactly one of data or filepath")
104        if data is None:
105            # filepath is not None from previous check
106            assert filepath is not None
107            data = Path(filepath).read_bytes()
108        name = _sha256_bytes(data) + suffix
109        bucket, key = self.parse_uri(name)
110        self.client.put_object(Bucket=bucket, Key=key, Body=data)
111        return Uri(f"s3://{bucket}/{key}")
112
113    def get(self, name_or_uri) -> bytes:
114        bucket, key = self.parse_uri(name_or_uri)
115        obj = self.client.get_object(Bucket=bucket, Key=key)
116        return obj["Body"].read()
117
118    def exists(self, name_or_uri) -> bool:
119        bucket, key = self.parse_uri(name_or_uri)
120        try:
121            self.client.head_object(Bucket=bucket, Key=key)
122            return True
123        except Exception as e:
124            code = getattr(e, "response", {}).get("Error", {}).get("Code")
125            if code in {"404", "NoSuchKey", "NotFound"}:
126                return False
127            raise
128
129    def ls(self, s3_root=None, *, recursive: bool = False, lazy: bool = False):
130        bucket, prefix = self.parse_uri(s3_root or self._name2uri(""))
131        if prefix:
132            prefix = prefix.rstrip("/") + "/"
133        kw: dict[str, Any] = {}
134        if not recursive:
135            kw["Delimiter"] = "/"
136        paginator = self.client.get_paginator("list_objects_v2")
137
138        def _iter():
139            for page in paginator.paginate(Bucket=bucket, Prefix=prefix, **kw):
140                for obj in page.get("Contents", []):
141                    yield Uri(f"s3://{bucket}/{obj['Key']}")
142
143        out = _iter()
144        if lazy:
145            return out
146        return list(out)
147
148    def rm(self, *name_or_uris):
149        values = _flatten_names(*name_or_uris)
150        if not values:
151            return
152        grouped: dict[str, list[str]] = {}
153        for item in values:
154            bucket, key = self.parse_uri(item)
155            grouped.setdefault(bucket, []).append(key)
156        for bucket, keys in grouped.items():
157            for i in range(0, len(keys), 1000):
158                batch = keys[i : i + 1000]
159                self.client.delete_objects(Bucket=bucket, Delete={"Objects": [{"Key": k} for k in batch]})
160
161    def put_js(self, data: Any) -> Uri:
162        encoded = json.dumps(data, separators=(",", ":"), sort_keys=True).encode("utf-8")
163        return self.put(data=encoded, suffix=".json")
164
165    def get_js(self, name_or_uri):
166        return json.loads(self.get(name_or_uri).decode("utf-8"))
167
168    def tar(
169        self,
170        path: str | os.PathLike[str],
171        excludes: Iterable[str] = (),
172        *,
173        symlinks: Literal["ignore", "raise"] = "raise",
174    ) -> Uri:
175        root = Path(path).resolve()
176        if not root.exists() or not root.is_dir():
177            raise DmlRepoError("S3Store.tar path must be an existing directory")
178        if symlinks not in {"ignore", "raise"}:
179            raise DmlRepoError("S3Store.tar symlinks must be 'ignore' or 'raise'")
180        patterns = list(excludes)
181        buf = io.BytesIO()
182
183        def excluded(rel: str) -> bool:
184            return any(fnmatch.fnmatch(rel, pat) for pat in patterns)
185
186        def normalize(info: tarfile.TarInfo) -> tarfile.TarInfo:
187            info.uid = 0
188            info.gid = 0
189            info.uname = ""
190            info.gname = ""
191            info.mtime = 0
192            return info
193
194        with tarfile.open(fileobj=buf, mode="w") as tf:
195            for dirpath, dirnames, filenames in os.walk(root):
196                dirpath = Path(dirpath)
197                rel_dir = dirpath.relative_to(root).as_posix()
198
199                kept_dirnames = []
200                for dirname in sorted(dirnames):
201                    child = dirpath / dirname
202                    rel = child.relative_to(root).as_posix()
203                    if excluded(rel):
204                        continue
205                    if child.is_symlink():
206                        if symlinks == "raise":
207                            raise DmlRepoError(f"S3Store.tar encountered symlink with symlinks='raise': {rel}")
208                        continue
209                    kept_dirnames.append(dirname)
210                dirnames[:] = kept_dirnames
211
212                if rel_dir != ".":
213                    tf.addfile(normalize(tf.gettarinfo(str(dirpath), arcname=rel_dir)))
214
215                for filename in sorted(filenames):
216                    p = dirpath / filename
217                    rel = p.relative_to(root).as_posix()
218                    if excluded(rel):
219                        continue
220                    if p.is_symlink():
221                        if symlinks == "raise":
222                            raise DmlRepoError(f"S3Store.tar encountered symlink with symlinks='raise': {rel}")
223                        continue
224                    with p.open("rb") as f:
225                        tf.addfile(normalize(tf.gettarinfo(str(p), arcname=rel)), fileobj=f)
226        return self.put(data=buf.getvalue(), suffix=".tar")
227
228    def untar(self, tar_uri, dest: str | os.PathLike[str], *, unsafe: bool = False) -> None:
229        payload = self.get(tar_uri)
230        dest_path = Path(dest)
231        dest_path.mkdir(parents=True, exist_ok=True)
232        resolved_dest = dest_path.resolve()
233        with tarfile.open(fileobj=io.BytesIO(payload), mode="r") as tf:
234            members = tf.getmembers()
235            if not unsafe:
236                for member in members:
237                    _validate_safe_extract_path(dest_path=resolved_dest, member_name=member.name)
238                tf.extractall(dest_path, members=members)
239                return
240            try:
241                tf.extractall(dest_path, members=members, filter="fully_trusted")
242            except TypeError:
243                tf.extractall(dest_path, members=members)
244
245    def cd(self, new_prefix: str) -> "S3Store":
246        current = Path("/" + self.prefix) if self.prefix else Path("/")
247        next_prefix = (current / new_prefix).resolve().as_posix().lstrip("/")
248        if next_prefix == ".":
249            next_prefix = ""
250        return S3Store(bucket=self.bucket, prefix=next_prefix, client=self.client)

is_s3_uri

def is_s3_uri(value: str) -> bool:
View source
20def is_s3_uri(value: str) -> bool:
21    p = urlparse(value)
22    return p.scheme == "s3" and bool(p.netloc) and bool(p.path and p.path != "/")

S3Store

@dataclass(frozen=True)
class S3Store:
View source
 44@dataclass(frozen=True)
 45class S3Store:
 46    bucket: str | None = None
 47    prefix: str | None = None
 48    client: Any = None
 49
 50    def __post_init__(self):
 51        bucket = self.bucket
 52        prefix = self.prefix
 53        if bucket is None and prefix is None:
 54            # use `daggerml`'s default dml session
 55            remote_root = get_default_dml().config.show()["remote"]["root"]
 56            if not remote_root:
 57                raise DmlRepoError(
 58                    "S3Store requires configured remote.root (set DML_REMOTE_ROOT or pass bucket/prefix)"
 59                )
 60            p = urlparse(remote_root)
 61            if p.scheme != "s3" or not p.netloc:
 62                raise DmlRepoError("remote.root must be an s3:// URI")
 63            bucket = p.netloc
 64            base = p.path.lstrip("/").rstrip("/")
 65            prefix = f"{base}/data" if base else "data"
 66        if bucket is None:
 67            raise DmlRepoError("S3Store bucket not configured")
 68        if prefix is None:
 69            prefix = ""
 70        object.__setattr__(self, "bucket", bucket)
 71        object.__setattr__(self, "prefix", prefix.strip("/"))
 72        object.__setattr__(self, "client", self.client or get_client("s3"))
 73
 74    @classmethod
 75    def from_remote_root(cls, remote_root: str) -> "S3Store":
 76        p = urlparse(remote_root)
 77        if p.scheme != "s3" or not p.netloc:
 78            raise DmlRepoError("remote root must be an s3:// URI")
 79        base = p.path.lstrip("/").rstrip("/")
 80        prefix = f"{base}/data" if base else "data"
 81        return cls(bucket=p.netloc, prefix=prefix)
 82
 83    def parse_uri(self, name_or_uri) -> tuple[str, str]:
 84        if isinstance(name_or_uri, Node):
 85            name_or_uri = name_or_uri.value()
 86        if isinstance(name_or_uri, Uri):
 87            name_or_uri = name_or_uri.uri
 88        if not isinstance(name_or_uri, str):
 89            raise DmlRepoError("S3Store name_or_uri must be a string or uri-bearing object")
 90        p = urlparse(name_or_uri)
 91        if p.scheme == "s3":
 92            return p.netloc, p.path[1:]
 93        if self.bucket is None:
 94            raise DmlRepoError("S3Store bucket not configured")
 95        key = f"{self.prefix}/{name_or_uri}" if self.prefix else name_or_uri
 96        return cast(str, self.bucket), key
 97
 98    def _name2uri(self, name) -> Uri:
 99        bucket, key = self.parse_uri(name)
100        return Uri(f"s3://{bucket}/{key}")
101
102    def put(self, data: bytes | None = None, filepath: str | None = None, *, suffix: str = "") -> Uri:
103        if (data is None) == (filepath is None):
104            raise DmlRepoError("S3Store.put requires exactly one of data or filepath")
105        if data is None:
106            # filepath is not None from previous check
107            assert filepath is not None
108            data = Path(filepath).read_bytes()
109        name = _sha256_bytes(data) + suffix
110        bucket, key = self.parse_uri(name)
111        self.client.put_object(Bucket=bucket, Key=key, Body=data)
112        return Uri(f"s3://{bucket}/{key}")
113
114    def get(self, name_or_uri) -> bytes:
115        bucket, key = self.parse_uri(name_or_uri)
116        obj = self.client.get_object(Bucket=bucket, Key=key)
117        return obj["Body"].read()
118
119    def exists(self, name_or_uri) -> bool:
120        bucket, key = self.parse_uri(name_or_uri)
121        try:
122            self.client.head_object(Bucket=bucket, Key=key)
123            return True
124        except Exception as e:
125            code = getattr(e, "response", {}).get("Error", {}).get("Code")
126            if code in {"404", "NoSuchKey", "NotFound"}:
127                return False
128            raise
129
130    def ls(self, s3_root=None, *, recursive: bool = False, lazy: bool = False):
131        bucket, prefix = self.parse_uri(s3_root or self._name2uri(""))
132        if prefix:
133            prefix = prefix.rstrip("/") + "/"
134        kw: dict[str, Any] = {}
135        if not recursive:
136            kw["Delimiter"] = "/"
137        paginator = self.client.get_paginator("list_objects_v2")
138
139        def _iter():
140            for page in paginator.paginate(Bucket=bucket, Prefix=prefix, **kw):
141                for obj in page.get("Contents", []):
142                    yield Uri(f"s3://{bucket}/{obj['Key']}")
143
144        out = _iter()
145        if lazy:
146            return out
147        return list(out)
148
149    def rm(self, *name_or_uris):
150        values = _flatten_names(*name_or_uris)
151        if not values:
152            return
153        grouped: dict[str, list[str]] = {}
154        for item in values:
155            bucket, key = self.parse_uri(item)
156            grouped.setdefault(bucket, []).append(key)
157        for bucket, keys in grouped.items():
158            for i in range(0, len(keys), 1000):
159                batch = keys[i : i + 1000]
160                self.client.delete_objects(Bucket=bucket, Delete={"Objects": [{"Key": k} for k in batch]})
161
162    def put_js(self, data: Any) -> Uri:
163        encoded = json.dumps(data, separators=(",", ":"), sort_keys=True).encode("utf-8")
164        return self.put(data=encoded, suffix=".json")
165
166    def get_js(self, name_or_uri):
167        return json.loads(self.get(name_or_uri).decode("utf-8"))
168
169    def tar(
170        self,
171        path: str | os.PathLike[str],
172        excludes: Iterable[str] = (),
173        *,
174        symlinks: Literal["ignore", "raise"] = "raise",
175    ) -> Uri:
176        root = Path(path).resolve()
177        if not root.exists() or not root.is_dir():
178            raise DmlRepoError("S3Store.tar path must be an existing directory")
179        if symlinks not in {"ignore", "raise"}:
180            raise DmlRepoError("S3Store.tar symlinks must be 'ignore' or 'raise'")
181        patterns = list(excludes)
182        buf = io.BytesIO()
183
184        def excluded(rel: str) -> bool:
185            return any(fnmatch.fnmatch(rel, pat) for pat in patterns)
186
187        def normalize(info: tarfile.TarInfo) -> tarfile.TarInfo:
188            info.uid = 0
189            info.gid = 0
190            info.uname = ""
191            info.gname = ""
192            info.mtime = 0
193            return info
194
195        with tarfile.open(fileobj=buf, mode="w") as tf:
196            for dirpath, dirnames, filenames in os.walk(root):
197                dirpath = Path(dirpath)
198                rel_dir = dirpath.relative_to(root).as_posix()
199
200                kept_dirnames = []
201                for dirname in sorted(dirnames):
202                    child = dirpath / dirname
203                    rel = child.relative_to(root).as_posix()
204                    if excluded(rel):
205                        continue
206                    if child.is_symlink():
207                        if symlinks == "raise":
208                            raise DmlRepoError(f"S3Store.tar encountered symlink with symlinks='raise': {rel}")
209                        continue
210                    kept_dirnames.append(dirname)
211                dirnames[:] = kept_dirnames
212
213                if rel_dir != ".":
214                    tf.addfile(normalize(tf.gettarinfo(str(dirpath), arcname=rel_dir)))
215
216                for filename in sorted(filenames):
217                    p = dirpath / filename
218                    rel = p.relative_to(root).as_posix()
219                    if excluded(rel):
220                        continue
221                    if p.is_symlink():
222                        if symlinks == "raise":
223                            raise DmlRepoError(f"S3Store.tar encountered symlink with symlinks='raise': {rel}")
224                        continue
225                    with p.open("rb") as f:
226                        tf.addfile(normalize(tf.gettarinfo(str(p), arcname=rel)), fileobj=f)
227        return self.put(data=buf.getvalue(), suffix=".tar")
228
229    def untar(self, tar_uri, dest: str | os.PathLike[str], *, unsafe: bool = False) -> None:
230        payload = self.get(tar_uri)
231        dest_path = Path(dest)
232        dest_path.mkdir(parents=True, exist_ok=True)
233        resolved_dest = dest_path.resolve()
234        with tarfile.open(fileobj=io.BytesIO(payload), mode="r") as tf:
235            members = tf.getmembers()
236            if not unsafe:
237                for member in members:
238                    _validate_safe_extract_path(dest_path=resolved_dest, member_name=member.name)
239                tf.extractall(dest_path, members=members)
240                return
241            try:
242                tf.extractall(dest_path, members=members, filter="fully_trusted")
243            except TypeError:
244                tf.extractall(dest_path, members=members)
245
246    def cd(self, new_prefix: str) -> "S3Store":
247        current = Path("/" + self.prefix) if self.prefix else Path("/")
248        next_prefix = (current / new_prefix).resolve().as_posix().lstrip("/")
249        if next_prefix == ".":
250            next_prefix = ""
251        return S3Store(bucket=self.bucket, prefix=next_prefix, client=self.client)

S3Store.__init__

S3Store( bucket: str | None = None, prefix: str | None = None, client: Any = None)

S3Store.bucket

bucket: str | None= None

S3Store.prefix

prefix: str | None= None

S3Store.client

client: Any= None

S3Store.from_remote_root

@classmethod
def from_remote_root(cls, remote_root: str) -> S3Store:
View source
74    @classmethod
75    def from_remote_root(cls, remote_root: str) -> "S3Store":
76        p = urlparse(remote_root)
77        if p.scheme != "s3" or not p.netloc:
78            raise DmlRepoError("remote root must be an s3:// URI")
79        base = p.path.lstrip("/").rstrip("/")
80        prefix = f"{base}/data" if base else "data"
81        return cls(bucket=p.netloc, prefix=prefix)

S3Store.parse_uri

def parse_uri(self, name_or_uri) -> tuple[str, str]:
View source
83    def parse_uri(self, name_or_uri) -> tuple[str, str]:
84        if isinstance(name_or_uri, Node):
85            name_or_uri = name_or_uri.value()
86        if isinstance(name_or_uri, Uri):
87            name_or_uri = name_or_uri.uri
88        if not isinstance(name_or_uri, str):
89            raise DmlRepoError("S3Store name_or_uri must be a string or uri-bearing object")
90        p = urlparse(name_or_uri)
91        if p.scheme == "s3":
92            return p.netloc, p.path[1:]
93        if self.bucket is None:
94            raise DmlRepoError("S3Store bucket not configured")
95        key = f"{self.prefix}/{name_or_uri}" if self.prefix else name_or_uri
96        return cast(str, self.bucket), key

S3Store.put

def put( self, data: bytes | None = None, filepath: str | None = None, *, suffix: str = '') -> daggerml.Uri:
View source
102    def put(self, data: bytes | None = None, filepath: str | None = None, *, suffix: str = "") -> Uri:
103        if (data is None) == (filepath is None):
104            raise DmlRepoError("S3Store.put requires exactly one of data or filepath")
105        if data is None:
106            # filepath is not None from previous check
107            assert filepath is not None
108            data = Path(filepath).read_bytes()
109        name = _sha256_bytes(data) + suffix
110        bucket, key = self.parse_uri(name)
111        self.client.put_object(Bucket=bucket, Key=key, Body=data)
112        return Uri(f"s3://{bucket}/{key}")

S3Store.get

def get(self, name_or_uri) -> bytes:
View source
114    def get(self, name_or_uri) -> bytes:
115        bucket, key = self.parse_uri(name_or_uri)
116        obj = self.client.get_object(Bucket=bucket, Key=key)
117        return obj["Body"].read()

S3Store.exists

def exists(self, name_or_uri) -> bool:
View source
119    def exists(self, name_or_uri) -> bool:
120        bucket, key = self.parse_uri(name_or_uri)
121        try:
122            self.client.head_object(Bucket=bucket, Key=key)
123            return True
124        except Exception as e:
125            code = getattr(e, "response", {}).get("Error", {}).get("Code")
126            if code in {"404", "NoSuchKey", "NotFound"}:
127                return False
128            raise

S3Store.ls

def ls(self, s3_root=None, *, recursive: bool = False, lazy: bool = False):
View source
130    def ls(self, s3_root=None, *, recursive: bool = False, lazy: bool = False):
131        bucket, prefix = self.parse_uri(s3_root or self._name2uri(""))
132        if prefix:
133            prefix = prefix.rstrip("/") + "/"
134        kw: dict[str, Any] = {}
135        if not recursive:
136            kw["Delimiter"] = "/"
137        paginator = self.client.get_paginator("list_objects_v2")
138
139        def _iter():
140            for page in paginator.paginate(Bucket=bucket, Prefix=prefix, **kw):
141                for obj in page.get("Contents", []):
142                    yield Uri(f"s3://{bucket}/{obj['Key']}")
143
144        out = _iter()
145        if lazy:
146            return out
147        return list(out)

S3Store.rm

def rm(self, *name_or_uris):
View source
149    def rm(self, *name_or_uris):
150        values = _flatten_names(*name_or_uris)
151        if not values:
152            return
153        grouped: dict[str, list[str]] = {}
154        for item in values:
155            bucket, key = self.parse_uri(item)
156            grouped.setdefault(bucket, []).append(key)
157        for bucket, keys in grouped.items():
158            for i in range(0, len(keys), 1000):
159                batch = keys[i : i + 1000]
160                self.client.delete_objects(Bucket=bucket, Delete={"Objects": [{"Key": k} for k in batch]})

S3Store.put_js

def put_js(self, data: Any) -> daggerml.Uri:
View source
162    def put_js(self, data: Any) -> Uri:
163        encoded = json.dumps(data, separators=(",", ":"), sort_keys=True).encode("utf-8")
164        return self.put(data=encoded, suffix=".json")

S3Store.get_js

def get_js(self, name_or_uri):
View source
166    def get_js(self, name_or_uri):
167        return json.loads(self.get(name_or_uri).decode("utf-8"))

S3Store.tar

def tar( self, path: str | os.PathLike[str], excludes: Iterable[str] = (), *, symlinks: Literal['ignore', 'raise'] = 'raise') -> daggerml.Uri:
View source
169    def tar(
170        self,
171        path: str | os.PathLike[str],
172        excludes: Iterable[str] = (),
173        *,
174        symlinks: Literal["ignore", "raise"] = "raise",
175    ) -> Uri:
176        root = Path(path).resolve()
177        if not root.exists() or not root.is_dir():
178            raise DmlRepoError("S3Store.tar path must be an existing directory")
179        if symlinks not in {"ignore", "raise"}:
180            raise DmlRepoError("S3Store.tar symlinks must be 'ignore' or 'raise'")
181        patterns = list(excludes)
182        buf = io.BytesIO()
183
184        def excluded(rel: str) -> bool:
185            return any(fnmatch.fnmatch(rel, pat) for pat in patterns)
186
187        def normalize(info: tarfile.TarInfo) -> tarfile.TarInfo:
188            info.uid = 0
189            info.gid = 0
190            info.uname = ""
191            info.gname = ""
192            info.mtime = 0
193            return info
194
195        with tarfile.open(fileobj=buf, mode="w") as tf:
196            for dirpath, dirnames, filenames in os.walk(root):
197                dirpath = Path(dirpath)
198                rel_dir = dirpath.relative_to(root).as_posix()
199
200                kept_dirnames = []
201                for dirname in sorted(dirnames):
202                    child = dirpath / dirname
203                    rel = child.relative_to(root).as_posix()
204                    if excluded(rel):
205                        continue
206                    if child.is_symlink():
207                        if symlinks == "raise":
208                            raise DmlRepoError(f"S3Store.tar encountered symlink with symlinks='raise': {rel}")
209                        continue
210                    kept_dirnames.append(dirname)
211                dirnames[:] = kept_dirnames
212
213                if rel_dir != ".":
214                    tf.addfile(normalize(tf.gettarinfo(str(dirpath), arcname=rel_dir)))
215
216                for filename in sorted(filenames):
217                    p = dirpath / filename
218                    rel = p.relative_to(root).as_posix()
219                    if excluded(rel):
220                        continue
221                    if p.is_symlink():
222                        if symlinks == "raise":
223                            raise DmlRepoError(f"S3Store.tar encountered symlink with symlinks='raise': {rel}")
224                        continue
225                    with p.open("rb") as f:
226                        tf.addfile(normalize(tf.gettarinfo(str(p), arcname=rel)), fileobj=f)
227        return self.put(data=buf.getvalue(), suffix=".tar")

S3Store.untar

def untar( self, tar_uri, dest: str | os.PathLike[str], *, unsafe: bool = False) -> None:
View source
229    def untar(self, tar_uri, dest: str | os.PathLike[str], *, unsafe: bool = False) -> None:
230        payload = self.get(tar_uri)
231        dest_path = Path(dest)
232        dest_path.mkdir(parents=True, exist_ok=True)
233        resolved_dest = dest_path.resolve()
234        with tarfile.open(fileobj=io.BytesIO(payload), mode="r") as tf:
235            members = tf.getmembers()
236            if not unsafe:
237                for member in members:
238                    _validate_safe_extract_path(dest_path=resolved_dest, member_name=member.name)
239                tf.extractall(dest_path, members=members)
240                return
241            try:
242                tf.extractall(dest_path, members=members, filter="fully_trusted")
243            except TypeError:
244                tf.extractall(dest_path, members=members)

S3Store.cd

def cd(self, new_prefix: str) -> S3Store:
View source
246    def cd(self, new_prefix: str) -> "S3Store":
247        current = Path("/" + self.prefix) if self.prefix else Path("/")
248        next_prefix = (current / new_prefix).resolve().as_posix().lstrip("/")
249        if next_prefix == ".":
250            next_prefix = ""
251        return S3Store(bucket=self.bucket, prefix=next_prefix, client=self.client)