sdrbench 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
sdrbench/__init__.py ADDED
@@ -0,0 +1,23 @@
1
+ """SDRBench scientific datasets as numpy arrays.
2
+
3
+ >>> import sdrbench
4
+ >>> sdrbench.list() # datasets; sdrbench.list(variants=True) for all variants
5
+ >>> nyx = sdrbench.dataset("nyx")
6
+ >>> nyx.fields
7
+ >>> t = nyx["temperature"] # memmap, shape (512, 512, 512), C order
8
+ >>> p = sdrbench.dataset("hurricane-isabel", "P").series("P") # 48 time steps
9
+
10
+ Files come from the Hugging Face mirror (https://huggingface.co/sdrbench), pinned to the
11
+ revision this release was built from; if Hugging Face cannot deliver, from the original
12
+ SDRBench archives on Globus. Local copies can be used via ``root=`` or ``$SDRBENCH_DATA``.
13
+ """
14
+ from ._core import Dataset, Field, Series, cache_dir, catalog, dataset, list, load
15
+
16
+ try:
17
+ from ._version import __version__
18
+ except ImportError: # source checkout without a build
19
+ __version__ = "0.0.0"
20
+
21
+ # `list` is deliberately not exported by `from sdrbench import *` (it would shadow the builtin);
22
+ # use sdrbench.list().
23
+ __all__ = ["Dataset", "Field", "Series", "cache_dir", "catalog", "dataset", "load"]
sdrbench/_archive.py ADDED
@@ -0,0 +1,219 @@
1
+ """Download and unpack SDRBench archives.
2
+
3
+ The Hugging Face mirror is built with exactly this code, so a file fetched
4
+ from Globus ends up at the same relative path, with the same bytes, as the
5
+ file stored on Hugging Face.
6
+ """
7
+ from __future__ import annotations
8
+
9
+ import hashlib
10
+ import http.client
11
+ import os
12
+ import shutil
13
+ import tarfile
14
+ import time
15
+ import urllib.error
16
+ import urllib.request
17
+ import zipfile
18
+ from pathlib import Path
19
+ from typing import Optional, Tuple
20
+
21
+ CHUNK = 1 << 22
22
+ JUNK = ("._", ".DS_Store", "__MACOSX")
23
+
24
+
25
+ def _is_junk(name: str) -> bool:
26
+ return any(part.startswith(JUNK) for part in Path(name).parts)
27
+
28
+
29
+ def sha256_file(path, chunk: int = 1 << 24) -> str:
30
+ h = hashlib.sha256()
31
+ with open(path, "rb") as f:
32
+ while b := f.read(chunk):
33
+ h.update(b)
34
+ return h.hexdigest()
35
+
36
+
37
+ class _HashingReader:
38
+ """File-like wrapper that hashes and counts everything read through it."""
39
+
40
+ def __init__(self, raw):
41
+ self.raw = raw
42
+ self.md5 = hashlib.md5()
43
+ self.nbytes = 0
44
+
45
+ def read(self, n: int = -1) -> bytes:
46
+ b = self.raw.read(n)
47
+ self.md5.update(b)
48
+ self.nbytes += len(b)
49
+ return b
50
+
51
+ def drain(self) -> None:
52
+ while self.read(CHUNK):
53
+ pass
54
+
55
+
56
+ class _ResumingResponse:
57
+ """Readable HTTP body that reconnects with a Range request when the connection drops,
58
+ so multi-GB downloads survive transient network failures.
59
+
60
+ The full size must be known (Globus sends GET bodies chunked, without Content-Length, so it
61
+ comes from HEAD); without it a dropped connection cannot be told apart from the end of the
62
+ file, so we refuse to download instead of risking a silently truncated archive."""
63
+
64
+ def __init__(self, url: str, retries: int = 10, timeout: float = 120):
65
+ self.url, self.retries, self.timeout = url, retries, timeout
66
+ self.pos = 0
67
+ self.total = self._head_length()
68
+ self.resp = self._connect()
69
+
70
+ def _head_length(self) -> int:
71
+ last = None
72
+ for attempt in range(self.retries + 1):
73
+ req = urllib.request.Request(self.url, method="HEAD", headers={"User-Agent": "sdrbench"})
74
+ try:
75
+ with urllib.request.urlopen(req, timeout=self.timeout) as r:
76
+ n = r.headers.get("Content-Length")
77
+ if n is not None:
78
+ return int(n)
79
+ last = "no Content-Length in HEAD response"
80
+ except urllib.error.HTTPError as e:
81
+ if e.code not in (429, 502, 503, 504): # only temporary server errors are retried
82
+ raise
83
+ last = e
84
+ except (OSError, http.client.HTTPException, ValueError) as e:
85
+ last = e
86
+ time.sleep(min(60, 2 ** attempt))
87
+ raise IOError(f"{self.url}: cannot determine the archive size ({last}); refusing an unverifiable download")
88
+
89
+ def _connect(self):
90
+ headers = {"User-Agent": "sdrbench"}
91
+ if self.pos:
92
+ headers["Range"] = f"bytes={self.pos}-"
93
+ resp = urllib.request.urlopen(urllib.request.Request(self.url, headers=headers), timeout=self.timeout)
94
+ if self.pos:
95
+ rng = resp.headers.get("Content-Range", "")
96
+ if resp.status != 206 or not rng.startswith(f"bytes {self.pos}-"):
97
+ resp.close()
98
+ raise IOError(f"{self.url}: server did not resume at byte {self.pos} ({resp.status} {rng!r})")
99
+ return resp
100
+
101
+ def read(self, n: int = -1) -> bytes:
102
+ for attempt in range(self.retries + 1):
103
+ try:
104
+ b = self.resp.read(n)
105
+ if not b and n != 0 and self.pos < self.total:
106
+ # body ended early: a dropped connection (with chunked encoding this can look
107
+ # like a clean end of stream)
108
+ raise http.client.IncompleteRead(b"", self.total - self.pos)
109
+ self.pos += len(b)
110
+ if self.pos > self.total:
111
+ raise IOError(f"{self.url}: received more than the {self.total} bytes announced")
112
+ return b
113
+ except (OSError, http.client.HTTPException) as e:
114
+ if isinstance(e, IOError) and "announced" in str(e) or attempt == self.retries:
115
+ raise
116
+ time.sleep(min(60, 2 ** attempt))
117
+ try:
118
+ self.resp.close()
119
+ except Exception:
120
+ pass
121
+ try:
122
+ self.resp = self._connect()
123
+ except (OSError, http.client.HTTPException):
124
+ continue # try again on the next attempt
125
+ raise IOError(f"{self.url}: download failed at byte {self.pos} of {self.total}")
126
+
127
+ def close(self):
128
+ self.resp.close()
129
+
130
+ def __enter__(self):
131
+ return self
132
+
133
+ def __exit__(self, *exc):
134
+ self.close()
135
+
136
+
137
+ def _open_url(url: str):
138
+ if url.startswith(("http://", "https://")):
139
+ return _ResumingResponse(url)
140
+ return urllib.request.urlopen(urllib.request.Request(url, headers={"User-Agent": "sdrbench"}), timeout=120)
141
+
142
+
143
+ def _safe_target(root: Path, name: str) -> Path:
144
+ target = (root / name).resolve()
145
+ if os.path.commonpath([str(root.resolve()), str(target)]) != str(root.resolve()):
146
+ raise ValueError(f"unsafe path in archive: {name}")
147
+ return target
148
+
149
+
150
+ def _extract_tar_stream(stream, root: Path) -> None:
151
+ with tarfile.open(fileobj=stream, mode="r|*") as tf:
152
+ for m in tf:
153
+ if _is_junk(m.name) or not (m.isfile() or m.isdir()):
154
+ continue
155
+ target = _safe_target(root, m.name)
156
+ if m.isdir():
157
+ target.mkdir(parents=True, exist_ok=True)
158
+ continue
159
+ target.parent.mkdir(parents=True, exist_ok=True)
160
+ src = tf.extractfile(m)
161
+ with open(target, "wb") as out:
162
+ shutil.copyfileobj(src, out, CHUNK)
163
+
164
+
165
+ def _extract_zip(path: Path, root: Path) -> None:
166
+ with zipfile.ZipFile(path) as zf:
167
+ for info in zf.infolist():
168
+ if _is_junk(info.filename) or info.is_dir():
169
+ continue
170
+ target = _safe_target(root, info.filename)
171
+ target.parent.mkdir(parents=True, exist_ok=True)
172
+ with zf.open(info) as src, open(target, "wb") as out:
173
+ shutil.copyfileobj(src, out, CHUNK)
174
+
175
+
176
+ def fetch_archive(url: str, dest, expected_md5: Optional[str] = None) -> Tuple[str, int]:
177
+ """Stream ``url`` (a .tar.gz or .zip) into directory ``dest``.
178
+
179
+ If all files of the archive live under a single top-level directory, its contents are
180
+ placed directly in ``dest`` (the same rule the Hugging Face mirror uses). File contents
181
+ are never modified. The whole archive must arrive: its size is checked against the
182
+ server's, and its md5 against ``expected_md5`` when given.
183
+ Returns ``(md5 of the archive, archive size in bytes)``.
184
+ """
185
+ import tempfile
186
+
187
+ dest = Path(dest)
188
+ dest.parent.mkdir(parents=True, exist_ok=True)
189
+ tmp = Path(tempfile.mkdtemp(prefix=dest.name + ".", suffix=".partial", dir=dest.parent))
190
+ try:
191
+ with _open_url(url) as resp:
192
+ reader = _HashingReader(resp)
193
+ if url.endswith(".zip"):
194
+ zpath = tmp / "_archive.zip"
195
+ with open(zpath, "wb") as out:
196
+ shutil.copyfileobj(reader, out, CHUNK)
197
+ files = tmp / "files"
198
+ _extract_zip(zpath, files)
199
+ zpath.unlink()
200
+ else:
201
+ files = tmp / "files"
202
+ files.mkdir()
203
+ _extract_tar_stream(reader, files)
204
+ reader.drain()
205
+ total = getattr(resp, "total", None)
206
+ if total is not None and reader.nbytes != total:
207
+ raise IOError(f"{url}: got {reader.nbytes} of {total} bytes")
208
+ md5 = reader.md5.hexdigest()
209
+ if expected_md5 and md5 != expected_md5:
210
+ raise IOError(f"md5 mismatch for {url}: got {md5}, expected {expected_md5}")
211
+ tops = {p.relative_to(files).parts[0] for p in files.rglob("*") if p.is_file()}
212
+ nested = {p.relative_to(files).parts[0] for p in files.rglob("*") if p.is_file()
213
+ and len(p.relative_to(files).parts) > 1}
214
+ src = files / next(iter(tops)) if len(tops) == 1 and tops == nested else files
215
+ shutil.rmtree(dest, ignore_errors=True)
216
+ os.replace(src, dest)
217
+ return md5, reader.nbytes
218
+ finally:
219
+ shutil.rmtree(tmp, ignore_errors=True)