908 lines
35 KiB
Python
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")
|