airfRANS-model-exploration/src/airfrans_frontier/raw/public.py

908 lines
35 KiB
Python

from __future__ import annotations
import json
import os
import shutil
import struct
import http.client
import time
import urllib.error
import urllib.request
import urllib.parse
import zlib
import zipfile
from dataclasses import dataclass, replace
from pathlib import Path, PurePosixPath
from typing import Any, Protocol
from airfrans_frontier.training.data_sources import publish_processed_dataset
PUBLIC_OF_DATASET_URL = "https://data.isir.upmc.fr/extrality/NeurIPS_2022/OF_dataset.zip"
DEFAULT_PUBLIC_WORK_DIR = Path("artifacts/public_airfrans")
DEFAULT_PUBLIC_OUTPUT_DIR = Path("artifacts/data_cache/airfrans_processed/processed/full")
_EOCD_SIGNATURE = b"PK\x05\x06"
_ZIP64_EOCD_LOCATOR_SIGNATURE = 0x07064B50
_ZIP64_EOCD_SIGNATURE = 0x06064B50
_CENTRAL_DIRECTORY_SIGNATURE = 0x02014B50
_LOCAL_FILE_HEADER_SIGNATURE = 0x04034B50
_ZIP64_EXTRA_ID = 0x0001
_ZIP64_LIMIT_16 = 0xFFFF
_ZIP64_LIMIT_32 = 0xFFFFFFFF
_HTTP_RANGE_READ_TIMEOUT_SECONDS = 60
_HTTP_RANGE_READ_MAX_ATTEMPTS = 4
_HTTP_RANGE_READ_RETRY_BASE_SECONDS = 2.0
class RangeReader(Protocol):
size: int
bytes_read: int
def read_range(self, start: int, length: int) -> bytes: ...
@dataclass(frozen=True)
class RemoteZipMember:
filename: str
flag_bits: int
compress_type: int
compress_size: int
file_size: int
header_offset: int
next_header_offset: int | None = None
@property
def is_dir(self) -> bool:
return self.filename.endswith("/")
@dataclass(frozen=True)
class StreamingZipProcessingResult:
processing: object
source_bytes: int
ranged_bytes_read: int
class PathRangeReader:
def __init__(self, path: str | Path) -> None:
self.path = Path(path).expanduser()
self.size = self.path.stat().st_size
self.bytes_read = 0
def read_range(self, start: int, length: int) -> bytes:
_validate_range(start, length, self.size)
if length == 0:
return b""
with self.path.open("rb") as handle:
handle.seek(start)
data = handle.read(length)
if len(data) != length:
raise RuntimeError(f"Local range read returned {len(data)} bytes; expected {length}")
self.bytes_read += len(data)
return data
class HttpRangeReader:
def __init__(self, url: str) -> None:
self.url = url
size = _remote_content_length(url)
if size is None:
raise RuntimeError(f"Could not determine remote content length for range streaming: {url}")
self.size = size
self.bytes_read = 0
def read_range(self, start: int, length: int) -> bytes:
_validate_range(start, length, self.size)
if length == 0:
return b""
end = start + length - 1
last_error: RuntimeError | None = None
for attempt in range(1, _HTTP_RANGE_READ_MAX_ATTEMPTS + 1):
request = urllib.request.Request(self.url, headers={"Range": f"bytes={start}-{end}"})
try:
with urllib.request.urlopen(request, timeout=_HTTP_RANGE_READ_TIMEOUT_SECONDS) as response:
status = getattr(response, "status", None)
data = response.read()
except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace")
last_error = RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc.code} {body}")
if exc.code in {408, 429} or 500 <= exc.code < 600:
if attempt < _HTTP_RANGE_READ_MAX_ATTEMPTS:
time.sleep(_HTTP_RANGE_READ_RETRY_BASE_SECONDS * attempt)
continue
raise last_error from exc
except (http.client.IncompleteRead, OSError) as exc:
last_error = RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc}")
if attempt < _HTTP_RANGE_READ_MAX_ATTEMPTS:
time.sleep(_HTTP_RANGE_READ_RETRY_BASE_SECONDS * attempt)
continue
raise last_error from exc
if status != 206:
raise RuntimeError(f"Server did not honor HTTP Range for {self.url}: status={status}")
if len(data) != length:
last_error = RuntimeError(f"HTTP range read returned {len(data)} bytes; expected {length}")
if attempt < _HTTP_RANGE_READ_MAX_ATTEMPTS:
time.sleep(_HTTP_RANGE_READ_RETRY_BASE_SECONDS * attempt)
continue
raise last_error
self.bytes_read += len(data)
return data
raise last_error or RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}")
def ensure_public_airfrans_processed_hf(
*,
repo_id: str,
path_in_repo: str = "processed/full",
work_dir: str | Path = DEFAULT_PUBLIC_WORK_DIR,
output_dir: str | Path = DEFAULT_PUBLIC_OUTPUT_DIR,
source_url: str = PUBLIC_OF_DATASET_URL,
min_cases: int = 1000,
private: bool = False,
force: bool = False,
) -> dict[str, Any]:
if min_cases <= 0:
raise ValueError("min_cases must be positive")
prefix = path_in_repo.strip("/")
started = time.time()
existing = _hf_dataset_status(repo_id=repo_id, path_in_repo=prefix)
if not force and existing["npz_file_count"] >= min_cases and existing["has_manifest"]:
return {
"ok": True,
"phase": "already_published",
"repo_id": repo_id,
"repo_type": "dataset",
"repo_url": f"https://huggingface.co/datasets/{repo_id}",
"path_in_repo": prefix,
"min_cases": min_cases,
"elapsed_seconds": time.time() - started,
**existing,
}
work_root = Path(work_dir).expanduser()
output_root = Path(output_dir).expanduser()
work_root.mkdir(parents=True, exist_ok=True)
output_root.mkdir(parents=True, exist_ok=True)
scratch_root = work_root / "streaming_raw"
print(f"range_stream_process_airfrans_zip source={source_url} output_dir={output_root}", flush=True)
streamed = process_of_dataset_url_streaming(
source_url,
output_root,
scratch_dir=scratch_root,
min_cases=min_cases,
force=force,
progress_every=25,
)
processed = streamed.processing
if processed.case_count < min_cases:
raise RuntimeError(f"Processed only {processed.case_count} cases from public AirfRANS archive; expected at least {min_cases}")
print(f"publish_airfrans_processed_hf repo={repo_id} path_in_repo={prefix}", flush=True)
publish = publish_processed_dataset(
data_root=output_root,
repo_id=repo_id,
path_in_repo=prefix,
private=private,
manifest_out=output_root / "hf_dataset_manifest.json",
)
final = _hf_dataset_status(repo_id=repo_id, path_in_repo=prefix)
if final["npz_file_count"] < min_cases:
raise RuntimeError(f"Published dataset has {final['npz_file_count']} .npz files under {prefix}; expected at least {min_cases}")
if not final["has_manifest"]:
raise RuntimeError(f"Published dataset is missing hf_dataset_manifest.json under {prefix}")
return {
"ok": True,
"phase": "published",
"repo_id": repo_id,
"repo_type": "dataset",
"repo_url": f"https://huggingface.co/datasets/{repo_id}",
"path_in_repo": prefix,
"source_url": source_url,
"streaming": True,
"streaming_mode": "zip_range",
"streaming_scratch_dir": str(scratch_root),
"source_bytes": streamed.source_bytes,
"ranged_bytes_read": streamed.ranged_bytes_read,
"output_dir": str(output_root),
"processed_case_count": processed.case_count,
"processed_total_points": processed.total_points,
"processed_manifest_path": str(processed.manifest_path),
"download": {"url": source_url, "mode": "zip_range", "source_bytes": streamed.source_bytes, "ranged_bytes_read": streamed.ranged_bytes_read},
"publish": publish,
"elapsed_seconds": time.time() - started,
**final,
}
def download_file(url: str, destination: str | Path, *, chunk_size: int = 16 * 1024 * 1024) -> dict[str, Any]:
path = Path(destination).expanduser()
path.parent.mkdir(parents=True, exist_ok=True)
expected_size = _remote_content_length(url)
existing_size = path.stat().st_size if path.exists() else 0
if expected_size is not None and existing_size == expected_size:
return {"url": url, "path": str(path), "bytes": existing_size, "resumed": False, "skipped": True}
headers: dict[str, str] = {}
mode = "wb"
resumed = False
if expected_size is not None and 0 < existing_size < expected_size:
headers["Range"] = f"bytes={existing_size}-"
mode = "ab"
resumed = True
print(
f"download_airfrans_zip url={url} path={path} existing_bytes={existing_size} expected_bytes={expected_size}",
flush=True,
)
request = urllib.request.Request(url, headers=headers)
try:
response = urllib.request.urlopen(request, timeout=60)
except urllib.error.HTTPError as exc:
if exc.code == 416 and expected_size is not None and existing_size >= expected_size:
return {"url": url, "path": str(path), "bytes": existing_size, "resumed": False, "skipped": True}
raise
with response:
if resumed and getattr(response, "status", None) != 206:
mode = "wb"
resumed = False
existing_size = 0
written = existing_size
next_report = ((written // 1_000_000_000) + 1) * 1_000_000_000
with path.open(mode) as handle:
while True:
chunk = response.read(chunk_size)
if not chunk:
break
handle.write(chunk)
written += len(chunk)
if written >= next_report:
print(f"downloaded_airfrans_zip_bytes={written}", flush=True)
next_report += 1_000_000_000
final_size = path.stat().st_size
if expected_size is not None and final_size != expected_size:
raise RuntimeError(f"Downloaded {final_size} bytes from {url}, expected {expected_size}")
return {"url": url, "path": str(path), "bytes": final_size, "resumed": resumed, "skipped": False}
def process_of_dataset_url_streaming(
source_url: str,
output_dir: str | Path,
*,
scratch_dir: str | Path,
min_cases: int = 1000,
force: bool = False,
progress_every: int | None = None,
) -> StreamingZipProcessingResult:
if min_cases <= 0:
raise ValueError("min_cases must be positive")
reader = _range_reader_for(source_url)
members = _read_zip_central_directory(reader)
processing = _process_remote_zip_members(
reader,
members,
output_dir,
scratch_dir=scratch_dir,
raw_dir_label=f"{source_url}!OF_dataset",
min_cases=min_cases,
force=force,
progress_every=progress_every,
)
print(
f"range_stream_airfrans_bytes_read={reader.bytes_read} range_stream_airfrans_source_bytes={reader.size}",
flush=True,
)
return StreamingZipProcessingResult(
processing=processing,
source_bytes=reader.size,
ranged_bytes_read=reader.bytes_read,
)
def _range_reader_for(source_url: str) -> RangeReader:
parsed = urllib.parse.urlparse(source_url)
if parsed.scheme in {"http", "https"}:
return HttpRangeReader(source_url)
if parsed.scheme == "file":
return PathRangeReader(Path(urllib.request.url2pathname(parsed.path)))
if not parsed.scheme:
return PathRangeReader(source_url)
raise RuntimeError(f"Unsupported AirfRANS streaming URL scheme: {parsed.scheme}")
def _read_zip_central_directory(reader: RangeReader) -> list[RemoteZipMember]:
tail_size = min(reader.size, 1024 * 1024)
tail_start = reader.size - tail_size
tail = reader.read_range(tail_start, tail_size)
eocd_index = tail.rfind(_EOCD_SIGNATURE)
if eocd_index < 0:
raise RuntimeError("ZIP end-of-central-directory record not found")
eocd_offset = tail_start + eocd_index
eocd = tail[eocd_index : eocd_index + 22]
if len(eocd) < 22:
raise RuntimeError("Truncated ZIP end-of-central-directory record")
(
_signature,
_disk_number,
_central_disk,
disk_entries,
total_entries,
central_size,
central_offset,
_comment_length,
) = struct.unpack("<IHHHHIIH", eocd)
if (
disk_entries == _ZIP64_LIMIT_16
or total_entries == _ZIP64_LIMIT_16
or central_size == _ZIP64_LIMIT_32
or central_offset == _ZIP64_LIMIT_32
):
total_entries, central_size, central_offset = _read_zip64_central_directory_locator(reader, eocd_offset)
central = reader.read_range(central_offset, central_size)
members = _annotate_next_header_offsets(_parse_central_directory(central, expected_entries=total_entries), central_offset=central_offset)
print(f"range_stream_airfrans_zip_members={len(members)}", flush=True)
return members
def _read_zip64_central_directory_locator(reader: RangeReader, eocd_offset: int) -> tuple[int, int, int]:
locator_offset = eocd_offset - 20
if locator_offset < 0:
raise RuntimeError("ZIP64 end-of-central-directory locator is missing")
locator = reader.read_range(locator_offset, 20)
signature, _disk_with_record, zip64_eocd_offset, _disk_count = struct.unpack("<IIQI", locator)
if signature != _ZIP64_EOCD_LOCATOR_SIGNATURE:
raise RuntimeError("ZIP64 end-of-central-directory locator has invalid signature")
record = reader.read_range(zip64_eocd_offset, 56)
(
record_signature,
_record_size,
_version_made,
_version_needed,
_disk_number,
_central_disk,
_disk_entries,
total_entries,
central_size,
central_offset,
) = struct.unpack("<IQHHIIQQQQ", record)
if record_signature != _ZIP64_EOCD_SIGNATURE:
raise RuntimeError("ZIP64 end-of-central-directory record has invalid signature")
return int(total_entries), int(central_size), int(central_offset)
def _parse_central_directory(central: bytes, *, expected_entries: int) -> list[RemoteZipMember]:
members: list[RemoteZipMember] = []
offset = 0
while offset < len(central):
if offset + 46 > len(central):
raise RuntimeError("Truncated ZIP central directory entry")
fields = struct.unpack_from("<IHHHHHHIIIHHHHHII", central, offset)
signature = fields[0]
if signature != _CENTRAL_DIRECTORY_SIGNATURE:
raise RuntimeError(f"Invalid ZIP central directory signature at offset {offset}")
flag_bits = fields[3]
compress_type = fields[4]
compress_size = fields[8]
file_size = fields[9]
filename_length = fields[10]
extra_length = fields[11]
comment_length = fields[12]
header_offset = fields[16]
name_start = offset + 46
extra_start = name_start + filename_length
comment_start = extra_start + extra_length
next_offset = comment_start + comment_length
if next_offset > len(central):
raise RuntimeError("Truncated ZIP central directory variable fields")
filename_bytes = central[name_start:extra_start]
encoding = "utf-8" if flag_bits & 0x800 else "cp437"
filename = filename_bytes.decode(encoding, errors="replace")
extra = central[extra_start:comment_start]
file_size, compress_size, header_offset = _apply_zip64_extra(
extra,
file_size=file_size,
compress_size=compress_size,
header_offset=header_offset,
)
members.append(
RemoteZipMember(
filename=filename,
flag_bits=flag_bits,
compress_type=compress_type,
compress_size=compress_size,
file_size=file_size,
header_offset=header_offset,
)
)
offset = next_offset
if expected_entries not in (0, len(members)):
raise RuntimeError(f"ZIP central directory entry count mismatch: parsed={len(members)} expected={expected_entries}")
return members
def _annotate_next_header_offsets(members: list[RemoteZipMember], *, central_offset: int) -> list[RemoteZipMember]:
next_offsets: dict[int, int] = {}
ordered = sorted(enumerate(members), key=lambda item: item[1].header_offset)
for position, (original_index, _member) in enumerate(ordered):
next_offsets[original_index] = (
ordered[position + 1][1].header_offset if position + 1 < len(ordered) else central_offset
)
return [replace(member, next_header_offset=next_offsets[index]) for index, member in enumerate(members)]
def _apply_zip64_extra(extra: bytes, *, file_size: int, compress_size: int, header_offset: int) -> tuple[int, int, int]:
values_needed = [
file_size == _ZIP64_LIMIT_32,
compress_size == _ZIP64_LIMIT_32,
header_offset == _ZIP64_LIMIT_32,
]
if not any(values_needed):
return file_size, compress_size, header_offset
offset = 0
while offset + 4 <= len(extra):
header_id, data_size = struct.unpack_from("<HH", extra, offset)
data_start = offset + 4
data_end = data_start + data_size
if data_end > len(extra):
raise RuntimeError("Truncated ZIP extra field")
if header_id == _ZIP64_EXTRA_ID:
cursor = data_start
resolved = [file_size, compress_size, header_offset]
for index, needed in enumerate(values_needed):
if needed:
if cursor + 8 > data_end:
raise RuntimeError("Truncated ZIP64 extra field")
resolved[index] = struct.unpack_from("<Q", extra, cursor)[0]
cursor += 8
return int(resolved[0]), int(resolved[1]), int(resolved[2])
offset = data_end
raise RuntimeError("ZIP64 central directory entry missing ZIP64 extra field")
def _process_remote_zip_members(
reader: RangeReader,
members: list[RemoteZipMember],
output_dir: str | Path,
*,
scratch_dir: str | Path,
raw_dir_label: str,
min_cases: int,
force: bool,
progress_every: int | None,
):
from airfrans_frontier.raw.process import process_raw_case_to_npz, write_processing_manifest
out_root = Path(output_dir).expanduser()
scratch_root = Path(scratch_dir).expanduser()
out_root.mkdir(parents=True, exist_ok=True)
if scratch_root.exists():
shutil.rmtree(scratch_root)
scratch_root.mkdir(parents=True, exist_ok=True)
case_members = _remote_archive_case_members(members)
case_names = sorted(case_members)
if len(case_names) < min_cases:
raise RuntimeError(f"AirfRANS archive has {len(case_names)} cases; expected at least {min_cases}")
print(f"range_stream_airfrans_archive_cases={len(case_names)}", flush=True)
records: list[dict[str, object]] = []
total_points = 0
started = time.perf_counter()
for index, case_name in enumerate(case_names, start=1):
case_dir = scratch_root / case_name
target_path = out_root / f"{case_name}.npz"
if target_path.exists() and not force:
record, points = process_raw_case_to_npz(case_dir, out_root, force=False)
else:
try:
_extract_remote_case_members(reader, case_members[case_name], scratch_root)
record, points = process_raw_case_to_npz(case_dir, out_root, force=force)
finally:
if case_dir.exists():
shutil.rmtree(case_dir, ignore_errors=True)
records.append(record)
total_points += points
if progress_every is not None and progress_every > 0 and (index % progress_every == 0 or index == len(case_names)):
print(f"range_streamed_airfrans_cases={index}/{len(case_names)} total_points={total_points}", flush=True)
try:
scratch_root.rmdir()
except OSError:
pass
return write_processing_manifest(
out_root,
raw_dir_label,
records=records,
total_points=total_points,
started=started,
)
def _remote_archive_case_members(members: list[RemoteZipMember]) -> dict[str, list[tuple[RemoteZipMember, PurePosixPath]]]:
cases: dict[str, list[tuple[RemoteZipMember, PurePosixPath]]] = {}
for member in members:
parsed = _case_member_parts_from_name(member.filename)
if parsed is None:
continue
case_name, relative = parsed
cases.setdefault(case_name, []).append((member, relative))
return cases
def _extract_remote_case_members(
reader: RangeReader,
members: list[tuple[RemoteZipMember, PurePosixPath]],
root: Path,
) -> None:
resolved_root = root.resolve()
entries = sorted(members, key=lambda item: item[0].header_offset)
span = _contiguous_case_span([member for member, _relative in entries])
if span is not None:
span_start, span_end = span
archive_bytes = reader.read_range(span_start, span_end - span_start)
for member, relative in entries:
target = _safe_relative_target(root, relative, resolved_root=resolved_root)
if member.is_dir:
target.mkdir(parents=True, exist_ok=True)
continue
target.parent.mkdir(parents=True, exist_ok=True)
target.write_bytes(_read_remote_member_payload_from_span(member, archive_bytes, span_start))
return
for member, relative in entries:
target = _safe_relative_target(root, relative, resolved_root=resolved_root)
if member.is_dir:
target.mkdir(parents=True, exist_ok=True)
continue
target.parent.mkdir(parents=True, exist_ok=True)
payload = _read_remote_member_payload(reader, member)
target.write_bytes(payload)
def _contiguous_case_span(members: list[RemoteZipMember]) -> tuple[int, int] | None:
if not members:
return None
ordered = sorted(members, key=lambda member: member.header_offset)
for current, following in zip(ordered, ordered[1:]):
if current.next_header_offset != following.header_offset:
return None
span_end = ordered[-1].next_header_offset
if span_end is None or span_end <= ordered[0].header_offset:
return None
return ordered[0].header_offset, span_end
def _read_remote_member_payload(reader: RangeReader, member: RemoteZipMember) -> bytes:
local_header = reader.read_range(member.header_offset, 30)
(
signature,
_version_needed,
_flag_bits,
_compress_type,
_mod_time,
_mod_date,
_crc,
_compress_size,
_file_size,
filename_length,
extra_length,
) = struct.unpack("<IHHHHHIIIHH", local_header)
if signature != _LOCAL_FILE_HEADER_SIGNATURE:
raise RuntimeError(f"Invalid local ZIP header for member: {member.filename}")
data_offset = member.header_offset + 30 + filename_length + extra_length
compressed = reader.read_range(data_offset, member.compress_size)
return _decode_remote_member_payload(member, compressed)
def _read_remote_member_payload_from_span(member: RemoteZipMember, archive_bytes: bytes, span_start: int) -> bytes:
local_header_offset = member.header_offset - span_start
local_header = archive_bytes[local_header_offset : local_header_offset + 30]
if len(local_header) != 30:
raise RuntimeError(f"Truncated local ZIP header for member: {member.filename}")
(
signature,
_version_needed,
_flag_bits,
_compress_type,
_mod_time,
_mod_date,
_crc,
_compress_size,
_file_size,
filename_length,
extra_length,
) = struct.unpack("<IHHHHHIIIHH", local_header)
if signature != _LOCAL_FILE_HEADER_SIGNATURE:
raise RuntimeError(f"Invalid local ZIP header for member: {member.filename}")
data_offset = local_header_offset + 30 + filename_length + extra_length
data_end = data_offset + member.compress_size
compressed = archive_bytes[data_offset:data_end]
if len(compressed) != member.compress_size:
raise RuntimeError(f"Truncated ZIP member payload for member: {member.filename}")
return _decode_remote_member_payload(member, compressed)
def _decode_remote_member_payload(member: RemoteZipMember, compressed: bytes) -> bytes:
if member.flag_bits & 0x1:
raise RuntimeError(f"Encrypted ZIP member is unsupported: {member.filename}")
if member.compress_type == 0:
payload = compressed
elif member.compress_type == 8:
decompressor = zlib.decompressobj(-15)
payload = decompressor.decompress(compressed) + decompressor.flush()
else:
raise RuntimeError(f"Unsupported ZIP compression method {member.compress_type} for {member.filename}")
if len(payload) != member.file_size:
raise RuntimeError(f"ZIP member size mismatch for {member.filename}: got {len(payload)} expected {member.file_size}")
return payload
def _validate_range(start: int, length: int, size: int) -> None:
if start < 0 or length < 0 or start + length > size:
raise RuntimeError(f"Invalid range start={start} length={length} size={size}")
def process_of_dataset_archive_streaming(
archive_path: str | Path,
output_dir: str | Path,
*,
scratch_dir: str | Path,
min_cases: int = 1000,
force: bool = False,
progress_every: int | None = None,
):
if min_cases <= 0:
raise ValueError("min_cases must be positive")
archive = Path(archive_path).expanduser()
out_root = Path(output_dir).expanduser()
scratch_root = Path(scratch_dir).expanduser()
out_root.mkdir(parents=True, exist_ok=True)
if scratch_root.exists():
shutil.rmtree(scratch_root)
scratch_root.mkdir(parents=True, exist_ok=True)
from airfrans_frontier.raw.process import process_raw_case_to_npz, write_processing_manifest
records: list[dict[str, object]] = []
total_points = 0
started = time.perf_counter()
with zipfile.ZipFile(archive) as zf:
case_members = _archive_case_members(zf.infolist())
case_names = sorted(case_members)
if len(case_names) < min_cases:
raise RuntimeError(f"AirfRANS archive has {len(case_names)} cases; expected at least {min_cases}")
print(f"stream_airfrans_archive_cases={len(case_names)}", flush=True)
for index, case_name in enumerate(case_names, start=1):
case_dir = scratch_root / case_name
target_path = out_root / f"{case_name}.npz"
if target_path.exists() and not force:
record, points = process_raw_case_to_npz(case_dir, out_root, force=False)
else:
try:
_extract_case_members(zf, case_members[case_name], scratch_root)
record, points = process_raw_case_to_npz(case_dir, out_root, force=force)
finally:
if case_dir.exists():
shutil.rmtree(case_dir, ignore_errors=True)
records.append(record)
total_points += points
if progress_every is not None and progress_every > 0 and (index % progress_every == 0 or index == len(case_names)):
print(f"streamed_airfrans_cases={index}/{len(case_names)} total_points={total_points}", flush=True)
try:
scratch_root.rmdir()
except OSError:
pass
return write_processing_manifest(
out_root,
f"{archive}!OF_dataset",
records=records,
total_points=total_points,
started=started,
)
def extract_of_dataset(archive_path: str | Path, extract_root: str | Path, *, min_cases: int = 1000) -> Path:
archive = Path(archive_path).expanduser()
root = Path(extract_root).expanduser()
root.mkdir(parents=True, exist_ok=True)
existing = _find_of_dataset_root(root)
if existing is not None and _case_count(existing) >= min_cases:
return existing
print(f"extract_airfrans_zip archive={archive} root={root}", flush=True)
with zipfile.ZipFile(archive) as zf:
members = zf.infolist()
_require_extract_space(root, members)
for index, member in enumerate(members, start=1):
_safe_extract_member(zf, member, root)
if index % 1000 == 0 or index == len(members):
print(f"extracted_airfrans_members={index}/{len(members)}", flush=True)
found = _find_of_dataset_root(root)
if found is None:
raise RuntimeError(f"OF_dataset directory not found after extracting {archive}")
case_count = _case_count(found)
if case_count < min_cases:
raise RuntimeError(f"Extracted AirfRANS OF_dataset has {case_count} cases; expected at least {min_cases}")
return found
def _hf_dataset_status(*, repo_id: str, path_in_repo: str) -> dict[str, Any]:
try:
from huggingface_hub import HfApi
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required for AirfRANS public data preparation") from exc
token = _optional_secret("HF_TOKEN")
api = HfApi(token=token)
try:
files = api.list_repo_files(repo_id=repo_id, repo_type="dataset")
except Exception:
files = []
prefix = path_in_repo.strip("/")
base = f"{prefix}/" if prefix else ""
npz_count = sum(1 for item in files if item.startswith(base) and item.endswith(".npz"))
has_manifest = any(item == f"{base}hf_dataset_manifest.json" for item in files)
return {
"file_count": len(files),
"npz_file_count": npz_count,
"has_manifest": has_manifest,
}
def _remote_content_length(url: str) -> int | None:
request = urllib.request.Request(url, method="HEAD")
try:
with urllib.request.urlopen(request, timeout=60) as response:
raw = response.headers.get("Content-Length")
except Exception:
return None
if raw is None:
return None
try:
return int(raw)
except ValueError:
return None
def _safe_extract_member(zf: zipfile.ZipFile, member: zipfile.ZipInfo, root: Path) -> None:
target = _safe_member_target(member, root)
if member.is_dir():
target.mkdir(parents=True, exist_ok=True)
return
target.parent.mkdir(parents=True, exist_ok=True)
with zf.open(member) as source, target.open("wb") as destination:
shutil.copyfileobj(source, destination, length=16 * 1024 * 1024)
def _require_extract_space(root: Path, members: list[zipfile.ZipInfo]) -> None:
total_uncompressed_bytes = 0
remaining_uncompressed_bytes = 0
resolved_root = root.resolve()
for member in members:
if member.is_dir():
continue
total_uncompressed_bytes += member.file_size
target = _safe_member_target(member, root, resolved_root=resolved_root)
try:
existing_size = target.stat().st_size
except OSError:
existing_size = None
if existing_size == member.file_size:
continue
remaining_uncompressed_bytes += member.file_size
margin_bytes = max(1024**3, remaining_uncompressed_bytes // 20) if remaining_uncompressed_bytes else 0
required_free_bytes = remaining_uncompressed_bytes + margin_bytes
usage = shutil.disk_usage(root)
print(
"airfrans_extract_total_uncompressed_bytes="
f"{total_uncompressed_bytes} airfrans_extract_remaining_uncompressed_bytes={remaining_uncompressed_bytes} "
f"airfrans_extract_free_disk_bytes={usage.free} airfrans_extract_required_free_bytes={required_free_bytes}",
flush=True,
)
if usage.free < required_free_bytes:
raise RuntimeError(
"Insufficient free disk for AirfRANS extraction: "
f"free={usage.free} required={required_free_bytes} remaining_uncompressed={remaining_uncompressed_bytes}; "
"provision more disk or use a streaming/incremental extraction pipeline"
)
def _safe_member_target(member: zipfile.ZipInfo, root: Path, *, resolved_root: Path | None = None) -> Path:
return _safe_relative_target(root, PurePosixPath(member.filename), resolved_root=resolved_root)
def _safe_relative_target(root: Path, relative: PurePosixPath, *, resolved_root: Path | None = None) -> Path:
target = root.joinpath(*relative.parts)
actual_root = resolved_root or root.resolve()
resolved_target = target.resolve()
if actual_root != resolved_target and actual_root not in resolved_target.parents:
raise RuntimeError(f"Unsafe path in AirfRANS archive: {relative}")
return target
def _archive_case_members(members: list[zipfile.ZipInfo]) -> dict[str, list[tuple[zipfile.ZipInfo, PurePosixPath]]]:
cases: dict[str, list[tuple[zipfile.ZipInfo, PurePosixPath]]] = {}
for member in members:
parsed = _case_member_parts(member)
if parsed is None:
continue
case_name, relative = parsed
cases.setdefault(case_name, []).append((member, relative))
return cases
def _case_member_parts(member: zipfile.ZipInfo) -> tuple[str, PurePosixPath] | None:
return _case_member_parts_from_name(member.filename)
def _case_member_parts_from_name(filename: str) -> tuple[str, PurePosixPath] | None:
parts = PurePosixPath(filename).parts
if any(part == ".." for part in parts):
raise RuntimeError(f"Unsafe path in AirfRANS archive: {filename}")
for index, part in enumerate(parts):
if part.startswith("airFoil2D_"):
return part, PurePosixPath(*parts[index:])
return None
def _extract_case_members(
zf: zipfile.ZipFile,
members: list[tuple[zipfile.ZipInfo, PurePosixPath]],
root: Path,
) -> None:
resolved_root = root.resolve()
for member, relative in members:
target = _safe_relative_target(root, relative, resolved_root=resolved_root)
if member.is_dir():
target.mkdir(parents=True, exist_ok=True)
continue
target.parent.mkdir(parents=True, exist_ok=True)
with zf.open(member) as source, target.open("wb") as destination:
shutil.copyfileobj(source, destination, length=16 * 1024 * 1024)
def _find_of_dataset_root(root: Path) -> Path | None:
direct = root / "OF_dataset"
if direct.is_dir():
return direct
for candidate in root.glob("*/OF_dataset"):
if candidate.is_dir():
return candidate
if _case_count(root) > 0:
return root
return None
def _case_count(root: Path) -> int:
return sum(1 for path in root.iterdir() if path.is_dir() and path.name.startswith("airFoil2D_")) if root.is_dir() else 0
def _optional_secret(name: str) -> str | None:
value = os.environ.get(name)
if value:
return value
for path in (Path(".env") / name, Path(".env") / f"{name}.txt"):
if path.is_file():
text = path.read_text().strip()
if text:
return text
return None
def _remove_file_best_effort(path: Path) -> bool:
try:
path.unlink()
return True
except FileNotFoundError:
return False
except OSError as exc:
print(f"warning: could not remove {path}: {exc}", flush=True)
return False
def write_json_report(path: str | Path, payload: dict[str, Any]) -> None:
report_path = Path(path).expanduser()
report_path.parent.mkdir(parents=True, exist_ok=True)
report_path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")