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

816 lines
31 KiB
Python
Raw Normal View History

2026-07-25 16:12:49 +00:00
from __future__ import annotations
import json
import os
import shutil
import struct
2026-07-25 16:12:49 +00:00
import time
import urllib.error
import urllib.request
import urllib.parse
import zlib
2026-07-25 16:12:49 +00:00
import zipfile
from dataclasses import dataclass
from pathlib import Path, PurePosixPath
from typing import Any, Protocol
2026-07-25 16:12:49 +00:00
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
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
@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
request = urllib.request.Request(self.url, headers={"Range": f"bytes={start}-{end}"})
try:
with urllib.request.urlopen(request, timeout=60) as response:
status = getattr(response, "status", None)
data = response.read()
except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace")
raise RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc.code} {body}") from exc
except OSError as exc:
raise RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc}") from exc
if status != 206:
raise RuntimeError(f"Server did not honor HTTP Range for {self.url}: status={status}")
if len(data) != length:
raise RuntimeError(f"HTTP range read returned {len(data)} bytes; expected {length}")
self.bytes_read += len(data)
return data
2026-07-25 16:12:49 +00:00
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
2026-07-25 16:12:49 +00:00
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,
2026-07-25 16:12:49 +00:00
"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},
2026-07-25 16:12:49 +00:00
"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 = _parse_central_directory(central, expected_entries=total_entries)
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 _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()
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)
payload = _read_remote_member_payload(reader, member)
target.write_bytes(payload)
def _read_remote_member_payload(reader: RangeReader, member: RemoteZipMember) -> bytes:
if member.flag_bits & 0x1:
raise RuntimeError(f"Encrypted ZIP member is unsupported: {member.filename}")
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)
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,
)
2026-07-25 16:12:49 +00:00
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)
2026-07-25 16:12:49 +00:00
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)
2026-07-25 16:12:49 +00:00
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)
2026-07-25 16:12:49 +00:00
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
2026-07-25 16:12:49 +00:00
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")