Nucleus-Resynthesis / runtime /src /resynthesis /direct_page_pack.py
Wl6adams's picture
Add portable Release 188 generation runtime
919fd68 verified
Raw
History Blame Contribute Delete
92 kB
"""Exact raw NoNE page-object pack sets.
This module is intentionally independent from ``none_paging`` so the storage
primitive can be validated before it is admitted into generation authority.
It never interprets revision-6 deltas and never compresses or reconstructs a
page object. Every packed entry is the byte-for-byte content of one complete
revision 2, 3, or 7 safetensors object.
The pack set is striped into disjoint shards. A single sequential device on
the target host is not fast enough for the all-knowledge cold-ingestion gate;
parallel shard reads provide one unique-byte measurement without counting
replica copies more than once.
"""
from __future__ import annotations
import hashlib
import mmap
import os
import stat
import time
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, replace
from pathlib import Path
from typing import Final, Sequence
import torch
from safetensors import safe_open
from safetensors.torch import save_file
DIRECT_PAGE_PACK_INDEX_REVISION: Final[int] = 1
DIRECT_PAGE_PACK_ALIGNMENT_BYTES: Final[int] = 4096
DIRECT_PAGE_PACK_IO_WAVE_BYTES: Final[int] = 64 * 1024 * 1024
DIRECT_PAGE_PACK_MAX_COALESCED_READ_BYTES: Final[int] = (
DIRECT_PAGE_PACK_IO_WAVE_BYTES
)
DIRECT_PAGE_PACK_MIN_BYTES_PER_SECOND: Final[int] = 1_000_000_000
DIRECT_PAGE_PACK_SUFFIX: Final[str] = ".none-direct-pack"
DIRECT_PAGE_PACK_INDEX_SUFFIX: Final[str] = ".none-direct-index.safetensors"
DIRECT_PAGE_PACK_FORMAT_REVISIONS: Final[frozenset[int]] = frozenset({2, 3, 7})
_PAGE_WEIGHT_TENSOR_NAMES: Final[tuple[str, ...]] = (
"ffn_mode_t",
"gate_t",
"up_t",
"down_t",
"glyph_down_t",
"glyph_up_t",
"translation_gate_t",
"outcome_memory_t",
"repair_memory_t",
"transfer_memory_t",
)
_REVISION_2_KEYS: Final[frozenset[str]] = frozenset(
{
"format_revision_t",
"page_ids_t",
"optimizer_mean_t",
"optimizer_square_t",
"step_t",
*_PAGE_WEIGHT_TENSOR_NAMES,
}
)
_REVISION_3_KEYS: Final[frozenset[str]] = frozenset(
{
"format_revision_t",
"page_ids_t",
"optimizer_width_t",
*_PAGE_WEIGHT_TENSOR_NAMES,
}
)
_REVISION_7_KEYS: Final[frozenset[str]] = frozenset(
{
"format_revision_t",
"page_ids_t",
"optimizer_width_t",
"step_t",
*_PAGE_WEIGHT_TENSOR_NAMES,
}
)
_FORBIDDEN_OBJECT_KEY_FRAGMENTS: Final[tuple[str, ...]] = (
"_delta_",
"base_object_",
"base_generation_",
"base_manifest_",
"compression",
"compressed",
"semantic_pack",
)
_INDEX_TENSOR_NAMES: Final[frozenset[str]] = frozenset(
{
"index_revision_t",
"alignment_bytes_t",
"page_count_t",
"shard_count_t",
"logical_object_bytes_t",
"physical_pack_bytes_t",
"page_ids_t",
"object_sha256s_t",
"object_bytes_t",
"shard_indices_t",
"object_offsets_t",
"object_spans_t",
"format_revisions_t",
"shard_sha256s_t",
"shard_bytes_t",
"shard_logical_bytes_t",
"shard_page_counts_t",
"page_ids_sha256_t",
"page_map_sha256_t",
"pack_set_sha256_t",
}
)
@dataclass(frozen=True)
class DirectPagePackSourcePacket:
"""External-I/O source rows for a deterministic raw pack build."""
page_ids_t: torch.Tensor
object_sha256s_t: torch.Tensor
object_bytes_t: torch.Tensor
object_paths: tuple[Path, ...]
@dataclass(frozen=True)
class DirectPagePackIndexPacket:
"""Tensor-native logical and physical index for one striped pack set."""
page_ids_t: torch.Tensor
object_sha256s_t: torch.Tensor
object_bytes_t: torch.Tensor
shard_indices_t: torch.Tensor
object_offsets_t: torch.Tensor
object_spans_t: torch.Tensor
format_revisions_t: torch.Tensor
shard_sha256s_t: torch.Tensor
shard_bytes_t: torch.Tensor
shard_logical_bytes_t: torch.Tensor
shard_page_counts_t: torch.Tensor
page_ids_sha256_t: torch.Tensor
page_map_sha256_t: torch.Tensor
pack_set_sha256_t: torch.Tensor
@dataclass(frozen=True)
class DirectPagePackSetAuthorityPacket:
"""Filesystem locator plus tensor identities for one immutable pack set."""
shard_roots: tuple[Path, ...]
shard_relative_paths: tuple[str, ...]
index_root: Path
index_relative_path: str
shard_sha256s_t: torch.Tensor
index_sha256_t: torch.Tensor
shard_bytes_t: torch.Tensor
index_bytes_t: torch.Tensor
logical_object_bytes_t: torch.Tensor
page_count_t: torch.Tensor
alignment_bytes_t: torch.Tensor
page_ids_sha256_t: torch.Tensor
page_map_sha256_t: torch.Tensor
pack_set_sha256_t: torch.Tensor
@dataclass(frozen=True)
class DirectPagePackBuildPacket:
"""Durable pack-set build result."""
authority: DirectPagePackSetAuthorityPacket
index: DirectPagePackIndexPacket
@dataclass(frozen=True)
class DirectPagePackSelectedPacket:
"""Exact selected raw objects in caller order."""
page_ids_t: torch.Tensor
object_sha256s_t: torch.Tensor
object_bytes_t: torch.Tensor
payload_offsets_t: torch.Tensor
object_payload_t: torch.Tensor
shard_physical_read_bytes_t: torch.Tensor
direct_io_t: torch.Tensor
@dataclass(frozen=True)
class DirectPagePackColdVerificationPacket:
"""Full unique-shard cold-ingestion evidence."""
pack_set_sha256_t: torch.Tensor
shard_sha256s_t: torch.Tensor
shard_physical_read_bytes_t: torch.Tensor
shard_logical_bytes_per_second_t: torch.Tensor
logical_object_bytes_t: torch.Tensor
physical_read_bytes_t: torch.Tensor
elapsed_nanoseconds_t: torch.Tensor
aggregate_unique_logical_bytes_per_second_t: torch.Tensor
direct_io_t: torch.Tensor
zero_padding_verified_t: torch.Tensor
@dataclass(frozen=True)
class DirectPagePackReplicaReceiptPacket:
"""One non-inflating durable replica-set result."""
authority: DirectPagePackSetAuthorityPacket
newly_written_shards_t: torch.Tensor
newly_written_logical_bytes_t: torch.Tensor
shard_replica_bytes_per_second_t: torch.Tensor
total_newly_written_logical_bytes_t: torch.Tensor
elapsed_nanoseconds_t: torch.Tensor
aggregate_unique_logical_bytes_per_second_t: torch.Tensor
direct_io_t: torch.Tensor
cold_verification: DirectPagePackColdVerificationPacket
@dataclass(frozen=True)
class DirectPagePackShardReplicaReceiptPacket:
"""One exact proof-only shard copy with zero duplicate speed credit."""
shard_index_t: torch.Tensor
shard_root: Path
shard_relative_path: str
shard_sha256_t: torch.Tensor
shard_bytes_t: torch.Tensor
logical_object_bytes_t: torch.Tensor
newly_written_t: torch.Tensor
newly_written_logical_bytes_t: torch.Tensor
elapsed_nanoseconds_t: torch.Tensor
direct_io_t: torch.Tensor
@dataclass(frozen=True)
class _BuiltShard:
shard_index: int
sha256: str
physical_bytes: int
logical_bytes: int
page_count: int
relative_path: str
def _stable_cpu_tensor(value_t: torch.Tensor, *, dtype: torch.dtype) -> torch.Tensor:
return value_t.detach().cpu().to(dtype=dtype).contiguous()
def _digest_tensor(raw_digest: bytes) -> torch.Tensor:
if len(raw_digest) != 32:
raise RuntimeError("direct page pack digest geometry differs")
return torch.frombuffer(bytearray(raw_digest), dtype=torch.uint8).clone()
def _tensor_digest(value_t: torch.Tensor) -> bytes:
stable_t = _stable_cpu_tensor(value_t, dtype=torch.uint8).reshape(-1)
if stable_t.shape != (32,):
raise RuntimeError("direct page pack digest tensor differs")
return stable_t.numpy().tobytes(order="C")
def _tensor_payload_digest(value_t: torch.Tensor) -> torch.Tensor:
stable_t = value_t.detach().cpu().contiguous()
digest = hashlib.sha256()
digest.update(str(stable_t.dtype).encode("ascii") + b"\x00")
digest.update(
torch.tensor(stable_t.shape, dtype=torch.long).numpy().tobytes(order="C")
)
digest.update(
stable_t.reshape(-1).view(torch.uint8).numpy().tobytes(order="C")
)
return _digest_tensor(digest.digest())
def _page_map_digest(
page_ids_t: torch.Tensor,
object_sha256s_t: torch.Tensor,
object_bytes_t: torch.Tensor,
) -> torch.Tensor:
digest = hashlib.sha256()
for name, value_t in (
("page_ids_t", page_ids_t),
("object_sha256s_t", object_sha256s_t),
("object_bytes_t", object_bytes_t),
):
stable_t = value_t.detach().cpu().contiguous()
digest.update(name.encode("ascii") + b"\x00")
digest.update(str(stable_t.dtype).encode("ascii") + b"\x00")
digest.update(
torch.tensor(stable_t.shape, dtype=torch.long)
.numpy()
.tobytes(order="C")
)
digest.update(
stable_t.reshape(-1).view(torch.uint8).numpy().tobytes(order="C")
)
return _digest_tensor(digest.digest())
def _pack_set_digest(index: DirectPagePackIndexPacket) -> torch.Tensor:
digest = hashlib.sha256()
for name, value_t in (
("page_ids_t", index.page_ids_t),
("object_sha256s_t", index.object_sha256s_t),
("object_bytes_t", index.object_bytes_t),
("shard_indices_t", index.shard_indices_t),
("object_offsets_t", index.object_offsets_t),
("object_spans_t", index.object_spans_t),
("format_revisions_t", index.format_revisions_t),
("shard_sha256s_t", index.shard_sha256s_t),
("shard_bytes_t", index.shard_bytes_t),
("shard_logical_bytes_t", index.shard_logical_bytes_t),
("shard_page_counts_t", index.shard_page_counts_t),
):
stable_t = value_t.detach().cpu().contiguous()
digest.update(name.encode("ascii") + b"\x00")
digest.update(str(stable_t.dtype).encode("ascii") + b"\x00")
digest.update(
torch.tensor(stable_t.shape, dtype=torch.long)
.numpy()
.tobytes(order="C")
)
digest.update(
stable_t.reshape(-1).view(torch.uint8).numpy().tobytes(order="C")
)
return _digest_tensor(digest.digest())
def _align_up(value: int, alignment: int = DIRECT_PAGE_PACK_ALIGNMENT_BYTES) -> int:
if value < 0 or alignment < 1 or alignment & (alignment - 1):
raise RuntimeError("direct page pack alignment differs")
return (value + alignment - 1) // alignment * alignment
def _file_identity(path: Path) -> tuple[int, int, int, int, int, int]:
identity = path.stat(follow_symlinks=False)
return (
identity.st_dev,
identity.st_ino,
identity.st_size,
identity.st_mtime_ns,
identity.st_ctime_ns,
identity.st_nlink,
)
def _file_sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
while True:
chunk = handle.read(8 * 1024 * 1024)
if not chunk:
break
digest.update(chunk)
return digest.hexdigest()
def _fsync_directory(path: Path) -> None:
descriptor = os.open(
path,
os.O_RDONLY
| getattr(os, "O_CLOEXEC", 0)
| getattr(os, "O_DIRECTORY", 0),
)
try:
os.fsync(descriptor)
finally:
os.close(descriptor)
def _reject_symlink_components(
path: Path,
*,
allow_missing_tail: bool,
) -> Path:
"""Reject unresolved symlinks in every existing path component."""
absolute = Path(os.path.abspath(os.fspath(path.expanduser())))
current = Path(absolute.anchor)
for part in absolute.parts[1:]:
current = current / part
try:
identity = current.lstat()
except FileNotFoundError:
if allow_missing_tail:
break
raise RuntimeError(
"direct page pack path component is absent"
) from None
if stat.S_ISLNK(identity.st_mode):
raise RuntimeError(
"direct page pack path contains a symlink component"
)
return absolute
def _safe_root(path: Path, *, create: bool) -> Path:
expanded = _reject_symlink_components(
path,
allow_missing_tail=create,
)
if create:
expanded.mkdir(parents=True, exist_ok=True)
_reject_symlink_components(expanded, allow_missing_tail=False)
resolved = expanded.resolve(strict=True)
identity = resolved.lstat()
if not stat.S_ISDIR(identity.st_mode):
raise RuntimeError("direct page pack storage root is not a directory")
return resolved
def _safe_relative_path(
root: Path,
relative_path: str,
*,
must_exist: bool,
create_parent: bool = False,
) -> Path:
relative = Path(relative_path)
if (
not relative_path
or relative.is_absolute()
or ".." in relative.parts
or relative.parts in {(), (".",)}
):
raise RuntimeError("direct page pack path escapes its storage root")
resolved_root = _safe_root(root, create=create_parent)
candidate = resolved_root / relative
if create_parent:
candidate.parent.mkdir(parents=True, exist_ok=True)
_reject_symlink_components(
candidate,
allow_missing_tail=not must_exist,
)
if must_exist:
resolved = candidate.resolve(strict=True)
if not resolved.is_relative_to(resolved_root):
raise RuntimeError("direct page pack path escapes its storage root")
identity = resolved.lstat()
if not stat.S_ISREG(identity.st_mode):
raise RuntimeError("direct page pack path is not a regular file")
return resolved
if not candidate.parent.resolve(strict=True).is_relative_to(resolved_root):
raise RuntimeError("direct page pack path escapes its storage root")
return candidate
def _source_object_identity(path: Path, *, expected_bytes: int) -> tuple[int, ...]:
expanded = _reject_symlink_components(
path,
allow_missing_tail=False,
)
resolved = expanded.resolve(strict=True)
identity = resolved.lstat()
if (
not stat.S_ISREG(identity.st_mode)
or identity.st_size != expected_bytes
or identity.st_nlink != 1
):
raise RuntimeError("direct page pack source object identity differs")
return _file_identity(resolved)
def _validate_source_object_schema(path: Path, *, expected_page_id: int) -> int:
with safe_open( # type: ignore[no-untyped-call]
str(path),
framework="pt",
device="cpu",
) as handle:
metadata = handle.metadata()
if metadata not in (None, {}):
raise RuntimeError(
"direct page pack source metadata must be empty"
)
keys = frozenset(str(name) for name in handle.keys())
if any(
fragment in name
for name in keys
for fragment in _FORBIDDEN_OBJECT_KEY_FRAGMENTS
):
raise RuntimeError(
"direct page pack source contains indirect or compressed state"
)
if "format_revision_t" not in keys or "page_ids_t" not in keys:
raise RuntimeError("direct page pack source scalar authority is absent")
revision_t = handle.get_tensor("format_revision_t").reshape(-1)
page_ids_t = handle.get_tensor("page_ids_t").reshape(-1)
if (
revision_t.dtype != torch.long
or revision_t.shape != (1,)
or page_ids_t.dtype != torch.long
or page_ids_t.shape != (1,)
):
raise RuntimeError("direct page pack source scalar authority differs")
revision = int(revision_t[0])
if revision not in DIRECT_PAGE_PACK_FORMAT_REVISIONS:
raise RuntimeError("direct page pack source revision is not 2, 3, or 7")
if int(page_ids_t[0]) != expected_page_id:
raise RuntimeError("direct page pack source page identity differs")
expected_keys = {
2: _REVISION_2_KEYS,
3: _REVISION_3_KEYS,
7: _REVISION_7_KEYS,
}[revision]
if keys != expected_keys:
raise RuntimeError("direct page pack source tensor key set differs")
weight_tensors = tuple(
handle.get_tensor(name) for name in _PAGE_WEIGHT_TENSOR_NAMES
)
storage_dtype = weight_tensors[0].dtype
if any(
tensor.dtype != storage_dtype
or tensor.ndim < 2
or tensor.shape[0] != 1
for tensor in weight_tensors
):
raise RuntimeError("direct page pack source weight geometry differs")
if revision in {3, 7}:
width_t = handle.get_tensor("optimizer_width_t").reshape(-1)
if (
width_t.dtype != torch.long
or width_t.shape != (1,)
or int(width_t[0]) < 1
):
raise RuntimeError(
"direct page pack source optimizer width differs"
)
if revision in {2, 7}:
step_t = handle.get_tensor("step_t").reshape(-1)
if step_t.dtype != torch.long or step_t.shape != (1,):
raise RuntimeError("direct page pack source optimizer step differs")
if revision == 2:
mean_t = handle.get_tensor("optimizer_mean_t")
square_t = handle.get_tensor("optimizer_square_t")
if (
mean_t.shape != square_t.shape
or mean_t.ndim != 2
or mean_t.shape[0] != 1
or mean_t.dtype not in {torch.float32, torch.bfloat16}
or square_t.dtype not in {torch.float32, torch.bfloat16}
):
raise RuntimeError(
"direct page pack source optimizer tensors differ"
)
return revision
def _validate_source_packet(
source: DirectPagePackSourcePacket,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[Path, ...]]:
page_ids_t = _stable_cpu_tensor(source.page_ids_t, dtype=torch.long).reshape(-1)
object_sha256s_t = _stable_cpu_tensor(
source.object_sha256s_t,
dtype=torch.uint8,
)
object_bytes_t = _stable_cpu_tensor(
source.object_bytes_t,
dtype=torch.long,
).reshape(-1)
count = int(page_ids_t.numel())
if (
source.page_ids_t.dtype != torch.long
or source.object_sha256s_t.dtype != torch.uint8
or source.object_bytes_t.dtype != torch.long
or count < 1
or object_sha256s_t.shape != (count, 32)
or object_bytes_t.shape != (count,)
or len(source.object_paths) != count
or bool(page_ids_t.lt(0).any())
or bool(object_bytes_t.le(0).any())
or torch.unique(page_ids_t).numel() != count
or torch.unique(object_sha256s_t, dim=0).shape[0] != count
):
raise RuntimeError("direct page pack source packet differs")
order_t = torch.argsort(page_ids_t, stable=True)
page_ids_t = page_ids_t[order_t].contiguous()
object_sha256s_t = object_sha256s_t[order_t].contiguous()
object_bytes_t = object_bytes_t[order_t].contiguous()
order = tuple(int(value) for value in order_t)
object_paths = tuple(source.object_paths[index] for index in order)
return page_ids_t, object_sha256s_t, object_bytes_t, object_paths
def _balanced_contiguous_shards(
object_spans_t: torch.Tensor,
*,
shard_count: int,
) -> torch.Tensor:
count = int(object_spans_t.numel())
if shard_count < 1 or shard_count > count:
raise ValueError("direct page pack shard count differs")
total_bytes = int(object_spans_t.sum())
target_bytes = (total_bytes + shard_count - 1) // shard_count
assignments = torch.empty(count, dtype=torch.long)
shard_index = 0
current_bytes = 0
for row_index in range(count):
span = int(object_spans_t[row_index])
remaining_rows = count - row_index
remaining_shards = shard_count - shard_index
if (
shard_index + 1 < shard_count
and current_bytes > 0
and current_bytes + span > target_bytes
and remaining_rows >= remaining_shards
):
shard_index += 1
current_bytes = 0
assignments[row_index] = shard_index
current_bytes += span
if int(assignments[-1]) != shard_count - 1:
# Very skewed object sizes can leave empty trailing shards. Move one
# terminal row into each missing shard without changing global order.
assignments = torch.arange(count, dtype=torch.long).mul(shard_count).div(
count,
rounding_mode="floor",
)
return assignments.contiguous()
def _write_source_into_shard(
*,
descriptor: int,
source_path: Path,
expected_page_id: int,
expected_sha256: bytes,
expected_bytes: int,
offset: int,
span: int,
) -> int:
identity_before = _source_object_identity(
source_path,
expected_bytes=expected_bytes,
)
resolved_source = source_path.expanduser().resolve(strict=True)
expected_hex = expected_sha256.hex()
if resolved_source.name != f"{expected_hex}.safetensors":
raise RuntimeError("direct page pack source content-addressed name differs")
revision = _validate_source_object_schema(
resolved_source,
expected_page_id=expected_page_id,
)
source_descriptor = os.open(
resolved_source,
os.O_RDONLY
| getattr(os, "O_CLOEXEC", 0)
| getattr(os, "O_NOFOLLOW", 0),
)
digest = hashlib.sha256()
copied = 0
try:
while copied < expected_bytes:
chunk = os.read(
source_descriptor,
min(8 * 1024 * 1024, expected_bytes - copied),
)
if not chunk:
raise RuntimeError("direct page pack source read was incomplete")
digest.update(chunk)
written = 0
while written < len(chunk):
count = os.pwrite(
descriptor,
chunk[written:],
offset + copied + written,
)
if count < 1:
raise RuntimeError(
"direct page pack destination write was incomplete"
)
written += count
copied += len(chunk)
finally:
os.close(source_descriptor)
if copied != expected_bytes or digest.digest() != expected_sha256:
raise RuntimeError("direct page pack source object hash differs")
padding_bytes = span - expected_bytes
if padding_bytes < 0 or padding_bytes >= DIRECT_PAGE_PACK_ALIGNMENT_BYTES:
raise RuntimeError("direct page pack source span differs")
if padding_bytes:
padding = b"\x00" * padding_bytes
written = os.pwrite(descriptor, padding, offset + expected_bytes)
if written != padding_bytes:
raise RuntimeError("direct page pack zero padding write was incomplete")
if _file_identity(resolved_source) != identity_before:
raise RuntimeError("direct page pack source changed during build")
return revision
def _build_one_shard(
*,
shard_index: int,
shard_root: Path,
row_indices: tuple[int, ...],
page_ids_t: torch.Tensor,
object_sha256s_t: torch.Tensor,
object_bytes_t: torch.Tensor,
object_paths: tuple[Path, ...],
object_offsets_t: torch.Tensor,
object_spans_t: torch.Tensor,
format_revisions_t: torch.Tensor,
) -> _BuiltShard:
root = _safe_root(shard_root, create=True)
temporary_root = _safe_relative_path(
root,
"direct-page-packs/sha256/.build",
must_exist=False,
create_parent=True,
)
temporary_root.mkdir(parents=True, exist_ok=True)
if temporary_root.is_symlink():
raise RuntimeError("direct page pack temporary root is a symlink")
temporary_path = temporary_root / (
f".shard-{shard_index}.{os.getpid()}.{time.monotonic_ns()}.tmp"
)
descriptor = -1
logical_bytes = 0
try:
descriptor = os.open(
temporary_path,
os.O_WRONLY
| os.O_CREAT
| os.O_EXCL
| getattr(os, "O_CLOEXEC", 0)
| getattr(os, "O_NOFOLLOW", 0),
0o600,
)
for row_index in row_indices:
object_sha256 = object_sha256s_t[row_index].numpy().tobytes()
revision = _write_source_into_shard(
descriptor=descriptor,
source_path=object_paths[row_index],
expected_page_id=int(page_ids_t[row_index]),
expected_sha256=object_sha256,
expected_bytes=int(object_bytes_t[row_index]),
offset=int(object_offsets_t[row_index]),
span=int(object_spans_t[row_index]),
)
format_revisions_t[row_index] = revision
logical_bytes += int(object_bytes_t[row_index])
physical_bytes = sum(int(object_spans_t[index]) for index in row_indices)
os.ftruncate(descriptor, physical_bytes)
os.fsync(descriptor)
os.close(descriptor)
descriptor = -1
shard_sha256 = _file_sha256(temporary_path)
relative_path = (
f"direct-page-packs/sha256/{shard_sha256}{DIRECT_PAGE_PACK_SUFFIX}"
)
final_path = _safe_relative_path(
root,
relative_path,
must_exist=False,
create_parent=True,
)
if final_path.exists():
if (
final_path.is_symlink()
or final_path.stat().st_size != temporary_path.stat().st_size
or _file_sha256(final_path) != shard_sha256
):
raise RuntimeError("existing direct page pack shard differs")
temporary_path.unlink()
else:
os.rename(temporary_path, final_path)
_fsync_directory(final_path.parent)
return _BuiltShard(
shard_index=shard_index,
sha256=shard_sha256,
physical_bytes=final_path.stat().st_size,
logical_bytes=logical_bytes,
page_count=len(row_indices),
relative_path=relative_path,
)
finally:
if descriptor >= 0:
os.close(descriptor)
temporary_path.unlink(missing_ok=True)
def _index_serialization_tensors(
index: DirectPagePackIndexPacket,
) -> dict[str, torch.Tensor]:
return {
"index_revision_t": torch.tensor(
[DIRECT_PAGE_PACK_INDEX_REVISION],
dtype=torch.long,
),
"alignment_bytes_t": torch.tensor(
[DIRECT_PAGE_PACK_ALIGNMENT_BYTES],
dtype=torch.long,
),
"page_count_t": torch.tensor(
[index.page_ids_t.numel()],
dtype=torch.long,
),
"shard_count_t": torch.tensor(
[index.shard_sha256s_t.shape[0]],
dtype=torch.long,
),
"logical_object_bytes_t": index.object_bytes_t.sum().reshape(1),
"physical_pack_bytes_t": index.shard_bytes_t.sum().reshape(1),
"page_ids_t": index.page_ids_t,
"object_sha256s_t": index.object_sha256s_t,
"object_bytes_t": index.object_bytes_t,
"shard_indices_t": index.shard_indices_t,
"object_offsets_t": index.object_offsets_t,
"object_spans_t": index.object_spans_t,
"format_revisions_t": index.format_revisions_t,
"shard_sha256s_t": index.shard_sha256s_t,
"shard_bytes_t": index.shard_bytes_t,
"shard_logical_bytes_t": index.shard_logical_bytes_t,
"shard_page_counts_t": index.shard_page_counts_t,
"page_ids_sha256_t": index.page_ids_sha256_t,
"page_map_sha256_t": index.page_map_sha256_t,
"pack_set_sha256_t": index.pack_set_sha256_t,
}
def validate_direct_page_pack_index_boundary(
index: DirectPagePackIndexPacket,
) -> DirectPagePackIndexPacket:
"""Fail closed on any noncanonical or indirect pack-set index."""
page_ids_t = _stable_cpu_tensor(index.page_ids_t, dtype=torch.long).reshape(-1)
object_sha256s_t = _stable_cpu_tensor(
index.object_sha256s_t,
dtype=torch.uint8,
)
object_bytes_t = _stable_cpu_tensor(
index.object_bytes_t,
dtype=torch.long,
).reshape(-1)
shard_indices_t = _stable_cpu_tensor(
index.shard_indices_t,
dtype=torch.long,
).reshape(-1)
object_offsets_t = _stable_cpu_tensor(
index.object_offsets_t,
dtype=torch.long,
).reshape(-1)
object_spans_t = _stable_cpu_tensor(
index.object_spans_t,
dtype=torch.long,
).reshape(-1)
format_revisions_t = _stable_cpu_tensor(
index.format_revisions_t,
dtype=torch.long,
).reshape(-1)
shard_sha256s_t = _stable_cpu_tensor(
index.shard_sha256s_t,
dtype=torch.uint8,
)
shard_bytes_t = _stable_cpu_tensor(
index.shard_bytes_t,
dtype=torch.long,
).reshape(-1)
shard_logical_bytes_t = _stable_cpu_tensor(
index.shard_logical_bytes_t,
dtype=torch.long,
).reshape(-1)
shard_page_counts_t = _stable_cpu_tensor(
index.shard_page_counts_t,
dtype=torch.long,
).reshape(-1)
count = int(page_ids_t.numel())
shard_count = int(shard_bytes_t.numel())
if (
index.page_ids_t.dtype != torch.long
or index.object_sha256s_t.dtype != torch.uint8
or index.object_bytes_t.dtype != torch.long
or index.shard_indices_t.dtype != torch.long
or index.object_offsets_t.dtype != torch.long
or index.object_spans_t.dtype != torch.long
or index.format_revisions_t.dtype != torch.long
or index.shard_sha256s_t.dtype != torch.uint8
or index.shard_bytes_t.dtype != torch.long
or index.shard_logical_bytes_t.dtype != torch.long
or index.shard_page_counts_t.dtype != torch.long
or count < 1
or shard_count < 1
or object_sha256s_t.shape != (count, 32)
or object_bytes_t.shape != (count,)
or shard_indices_t.shape != (count,)
or object_offsets_t.shape != (count,)
or object_spans_t.shape != (count,)
or format_revisions_t.shape != (count,)
or shard_sha256s_t.shape != (shard_count, 32)
or shard_logical_bytes_t.shape != (shard_count,)
or shard_page_counts_t.shape != (shard_count,)
or not torch.equal(page_ids_t, torch.sort(page_ids_t).values)
or torch.unique(page_ids_t).numel() != count
or torch.unique(object_sha256s_t, dim=0).shape[0] != count
or bool(object_bytes_t.le(0).any())
or bool(shard_indices_t.lt(0).any())
or bool(shard_indices_t.ge(shard_count).any())
or bool(object_offsets_t.lt(0).any())
or bool(object_spans_t.le(0).any())
or bool(object_spans_t.remainder(DIRECT_PAGE_PACK_ALIGNMENT_BYTES).any())
or bool(object_offsets_t.remainder(DIRECT_PAGE_PACK_ALIGNMENT_BYTES).any())
or bool(object_spans_t.lt(object_bytes_t).any())
or bool(
object_spans_t
.ne(
object_bytes_t.add(DIRECT_PAGE_PACK_ALIGNMENT_BYTES - 1)
.div(DIRECT_PAGE_PACK_ALIGNMENT_BYTES, rounding_mode="floor")
.mul(DIRECT_PAGE_PACK_ALIGNMENT_BYTES)
)
.any()
)
or bool(
~torch.isin(
format_revisions_t,
torch.tensor(
sorted(DIRECT_PAGE_PACK_FORMAT_REVISIONS),
dtype=torch.long,
),
).all()
)
or bool(shard_bytes_t.le(0).any())
or bool(shard_logical_bytes_t.le(0).any())
or bool(shard_page_counts_t.le(0).any())
or int(shard_page_counts_t.sum()) != count
):
raise RuntimeError("direct page pack index tensor authority differs")
derived_shard_bytes = torch.zeros(shard_count, dtype=torch.long)
derived_logical_bytes = torch.zeros(shard_count, dtype=torch.long)
derived_page_counts = torch.zeros(shard_count, dtype=torch.long)
for shard_index in range(shard_count):
rows_t = shard_indices_t.eq(shard_index).nonzero(
as_tuple=False
).reshape(-1)
if rows_t.numel() < 1:
raise RuntimeError("direct page pack index contains an empty shard")
offsets_t = object_offsets_t[rows_t]
spans_t = object_spans_t[rows_t]
expected_offsets_t = torch.cat(
(
torch.zeros(1, dtype=torch.long),
torch.cumsum(spans_t[:-1], dim=0),
)
)
if not torch.equal(offsets_t, expected_offsets_t):
raise RuntimeError(
"direct page pack index overlaps or contains a physical hole"
)
derived_shard_bytes[shard_index] = spans_t.sum()
derived_logical_bytes[shard_index] = object_bytes_t[rows_t].sum()
derived_page_counts[shard_index] = rows_t.numel()
if (
not torch.equal(derived_shard_bytes, shard_bytes_t)
or not torch.equal(derived_logical_bytes, shard_logical_bytes_t)
or not torch.equal(derived_page_counts, shard_page_counts_t)
):
raise RuntimeError("direct page pack shard aggregates differ")
canonical = DirectPagePackIndexPacket(
page_ids_t=page_ids_t,
object_sha256s_t=object_sha256s_t,
object_bytes_t=object_bytes_t,
shard_indices_t=shard_indices_t,
object_offsets_t=object_offsets_t,
object_spans_t=object_spans_t,
format_revisions_t=format_revisions_t,
shard_sha256s_t=shard_sha256s_t,
shard_bytes_t=shard_bytes_t,
shard_logical_bytes_t=shard_logical_bytes_t,
shard_page_counts_t=shard_page_counts_t,
page_ids_sha256_t=_stable_cpu_tensor(
index.page_ids_sha256_t,
dtype=torch.uint8,
).reshape(-1),
page_map_sha256_t=_stable_cpu_tensor(
index.page_map_sha256_t,
dtype=torch.uint8,
).reshape(-1),
pack_set_sha256_t=_stable_cpu_tensor(
index.pack_set_sha256_t,
dtype=torch.uint8,
).reshape(-1),
)
expected_page_ids_sha256_t = _tensor_payload_digest(page_ids_t)
expected_page_map_sha256_t = _page_map_digest(
page_ids_t,
object_sha256s_t,
object_bytes_t,
)
expected_pack_set_sha256_t = _pack_set_digest(
replace(
canonical,
page_ids_sha256_t=expected_page_ids_sha256_t,
page_map_sha256_t=expected_page_map_sha256_t,
pack_set_sha256_t=torch.zeros(32, dtype=torch.uint8),
)
)
if (
canonical.page_ids_sha256_t.shape != (32,)
or canonical.page_map_sha256_t.shape != (32,)
or canonical.pack_set_sha256_t.shape != (32,)
or not torch.equal(
canonical.page_ids_sha256_t,
expected_page_ids_sha256_t,
)
or not torch.equal(
canonical.page_map_sha256_t,
expected_page_map_sha256_t,
)
or not torch.equal(
canonical.pack_set_sha256_t,
expected_pack_set_sha256_t,
)
):
raise RuntimeError("direct page pack index digest authority differs")
return canonical
def build_direct_page_pack_set_boundary(
source: DirectPagePackSourcePacket,
*,
shard_roots: Sequence[Path],
index_root: Path,
) -> DirectPagePackBuildPacket:
"""Build deterministic disjoint aligned raw shards without whole-pack memory."""
(
page_ids_t,
object_sha256s_t,
object_bytes_t,
object_paths,
) = _validate_source_packet(source)
roots = tuple(_safe_root(Path(root), create=True) for root in shard_roots)
if not roots or len(roots) > page_ids_t.numel():
raise ValueError("direct page pack shard roots differ")
object_spans_t = (
object_bytes_t.add(DIRECT_PAGE_PACK_ALIGNMENT_BYTES - 1)
.div(DIRECT_PAGE_PACK_ALIGNMENT_BYTES, rounding_mode="floor")
.mul(DIRECT_PAGE_PACK_ALIGNMENT_BYTES)
.contiguous()
)
shard_indices_t = _balanced_contiguous_shards(
object_spans_t,
shard_count=len(roots),
)
object_offsets_t = torch.empty_like(object_spans_t)
for shard_index in range(len(roots)):
rows_t = shard_indices_t.eq(shard_index).nonzero(
as_tuple=False
).reshape(-1)
spans_t = object_spans_t[rows_t]
object_offsets_t[rows_t] = torch.cat(
(
torch.zeros(1, dtype=torch.long),
torch.cumsum(spans_t[:-1], dim=0),
)
)
format_revisions_t = torch.zeros_like(page_ids_t)
futures = []
with ThreadPoolExecutor(
max_workers=len(roots),
thread_name_prefix="nnf-direct-pack-build",
) as executor:
for shard_index, root in enumerate(roots):
row_indices = tuple(
int(value)
for value in shard_indices_t.eq(shard_index)
.nonzero(as_tuple=False)
.reshape(-1)
)
futures.append(
executor.submit(
_build_one_shard,
shard_index=shard_index,
shard_root=root,
row_indices=row_indices,
page_ids_t=page_ids_t,
object_sha256s_t=object_sha256s_t,
object_bytes_t=object_bytes_t,
object_paths=object_paths,
object_offsets_t=object_offsets_t,
object_spans_t=object_spans_t,
format_revisions_t=format_revisions_t,
)
)
built_shards = tuple(future.result() for future in futures)
built_shards = tuple(
sorted(built_shards, key=lambda shard: shard.shard_index)
)
shard_sha256s_t = torch.stack(
tuple(
_digest_tensor(bytes.fromhex(shard.sha256))
for shard in built_shards
)
)
shard_bytes_t = torch.tensor(
[shard.physical_bytes for shard in built_shards],
dtype=torch.long,
)
shard_logical_bytes_t = torch.tensor(
[shard.logical_bytes for shard in built_shards],
dtype=torch.long,
)
shard_page_counts_t = torch.tensor(
[shard.page_count for shard in built_shards],
dtype=torch.long,
)
page_ids_sha256_t = _tensor_payload_digest(page_ids_t)
page_map_sha256_t = _page_map_digest(
page_ids_t,
object_sha256s_t,
object_bytes_t,
)
provisional = DirectPagePackIndexPacket(
page_ids_t=page_ids_t,
object_sha256s_t=object_sha256s_t,
object_bytes_t=object_bytes_t,
shard_indices_t=shard_indices_t,
object_offsets_t=object_offsets_t,
object_spans_t=object_spans_t,
format_revisions_t=format_revisions_t,
shard_sha256s_t=shard_sha256s_t,
shard_bytes_t=shard_bytes_t,
shard_logical_bytes_t=shard_logical_bytes_t,
shard_page_counts_t=shard_page_counts_t,
page_ids_sha256_t=page_ids_sha256_t,
page_map_sha256_t=page_map_sha256_t,
pack_set_sha256_t=torch.zeros(32, dtype=torch.uint8),
)
index = replace(
provisional,
pack_set_sha256_t=_pack_set_digest(provisional),
)
index = validate_direct_page_pack_index_boundary(index)
resolved_index_root = _safe_root(index_root, create=True)
temporary_index = _safe_relative_path(
resolved_index_root,
(
"direct-page-packs/index/"
f".index.{os.getpid()}.{time.monotonic_ns()}.tmp"
),
must_exist=False,
create_parent=True,
)
save_file(_index_serialization_tensors(index), str(temporary_index))
with temporary_index.open("rb") as handle:
os.fsync(handle.fileno())
index_sha256 = _file_sha256(temporary_index)
index_relative_path = (
f"direct-page-packs/index/{index_sha256}"
f"{DIRECT_PAGE_PACK_INDEX_SUFFIX}"
)
final_index = _safe_relative_path(
resolved_index_root,
index_relative_path,
must_exist=False,
create_parent=True,
)
try:
if final_index.exists():
if (
final_index.is_symlink()
or final_index.stat().st_size != temporary_index.stat().st_size
or _file_sha256(final_index) != index_sha256
):
raise RuntimeError("existing direct page pack index differs")
temporary_index.unlink()
else:
os.rename(temporary_index, final_index)
_fsync_directory(final_index.parent)
finally:
temporary_index.unlink(missing_ok=True)
authority = DirectPagePackSetAuthorityPacket(
shard_roots=roots,
shard_relative_paths=tuple(
shard.relative_path for shard in built_shards
),
index_root=resolved_index_root,
index_relative_path=index_relative_path,
shard_sha256s_t=index.shard_sha256s_t.clone(),
index_sha256_t=_digest_tensor(bytes.fromhex(index_sha256)),
shard_bytes_t=index.shard_bytes_t.clone(),
index_bytes_t=torch.tensor(
final_index.stat().st_size,
dtype=torch.long,
),
logical_object_bytes_t=index.object_bytes_t.sum().reshape(()),
page_count_t=torch.tensor(index.page_ids_t.numel(), dtype=torch.long),
alignment_bytes_t=torch.tensor(
DIRECT_PAGE_PACK_ALIGNMENT_BYTES,
dtype=torch.long,
),
page_ids_sha256_t=index.page_ids_sha256_t.clone(),
page_map_sha256_t=index.page_map_sha256_t.clone(),
pack_set_sha256_t=index.pack_set_sha256_t.clone(),
)
validate_direct_page_pack_set_authority_boundary(authority, index=index)
return DirectPagePackBuildPacket(authority=authority, index=index)
def _load_index_file(path: Path) -> DirectPagePackIndexPacket:
with safe_open( # type: ignore[no-untyped-call]
str(path),
framework="pt",
device="cpu",
) as handle:
keys = frozenset(str(name) for name in handle.keys())
if keys != _INDEX_TENSOR_NAMES:
raise RuntimeError("direct page pack index key set differs")
scalars = {
name: handle.get_tensor(name).reshape(-1)
for name in (
"index_revision_t",
"alignment_bytes_t",
"page_count_t",
"shard_count_t",
"logical_object_bytes_t",
"physical_pack_bytes_t",
)
}
if any(
value.dtype != torch.long or value.shape != (1,)
for value in scalars.values()
):
raise RuntimeError("direct page pack index scalar authority differs")
if (
int(scalars["index_revision_t"][0])
!= DIRECT_PAGE_PACK_INDEX_REVISION
or int(scalars["alignment_bytes_t"][0])
!= DIRECT_PAGE_PACK_ALIGNMENT_BYTES
):
raise RuntimeError("direct page pack index revision differs")
index = DirectPagePackIndexPacket(
page_ids_t=handle.get_tensor("page_ids_t"),
object_sha256s_t=handle.get_tensor("object_sha256s_t"),
object_bytes_t=handle.get_tensor("object_bytes_t"),
shard_indices_t=handle.get_tensor("shard_indices_t"),
object_offsets_t=handle.get_tensor("object_offsets_t"),
object_spans_t=handle.get_tensor("object_spans_t"),
format_revisions_t=handle.get_tensor("format_revisions_t"),
shard_sha256s_t=handle.get_tensor("shard_sha256s_t"),
shard_bytes_t=handle.get_tensor("shard_bytes_t"),
shard_logical_bytes_t=handle.get_tensor(
"shard_logical_bytes_t"
),
shard_page_counts_t=handle.get_tensor("shard_page_counts_t"),
page_ids_sha256_t=handle.get_tensor("page_ids_sha256_t"),
page_map_sha256_t=handle.get_tensor("page_map_sha256_t"),
pack_set_sha256_t=handle.get_tensor("pack_set_sha256_t"),
)
validated = validate_direct_page_pack_index_boundary(index)
if (
int(scalars["page_count_t"][0]) != validated.page_ids_t.numel()
or int(scalars["shard_count_t"][0])
!= validated.shard_sha256s_t.shape[0]
or int(scalars["logical_object_bytes_t"][0])
!= int(validated.object_bytes_t.sum())
or int(scalars["physical_pack_bytes_t"][0])
!= int(validated.shard_bytes_t.sum())
):
raise RuntimeError("direct page pack index aggregate scalar differs")
return validated
def reopen_direct_page_pack_set_boundary(
*,
storage_root: Path,
index_relative_path: str,
) -> DirectPagePackBuildPacket:
"""Reopen one exact canonical pack set without rebuilding or copying it.
The caller supplies the immutable storage root and the exact
content-addressed index path. The index is the sole logical authority:
shard paths are derived from its shard digests, never discovered by
walking source objects or storage directories. This boundary reads and
validates the small index file, but it deliberately does not open or hash
any shard payload. Full payload hashing remains the responsibility of the
explicit cold-verification boundary.
"""
resolved_root = _safe_root(storage_root, create=False)
index_path = _safe_relative_path(
resolved_root,
index_relative_path,
must_exist=True,
)
index_identity = _file_identity(index_path)
if index_identity[5] != 1:
raise RuntimeError("direct page pack index file authority differs")
index_sha256 = _file_sha256(index_path)
expected_index_relative_path = (
"direct-page-packs/index/"
f"{index_sha256}{DIRECT_PAGE_PACK_INDEX_SUFFIX}"
)
if index_relative_path != expected_index_relative_path:
raise RuntimeError(
"direct page pack index content-addressed path differs"
)
if _file_identity(index_path) != index_identity:
raise RuntimeError("direct page pack index changed during validation")
index = _load_index_file(index_path)
if _file_identity(index_path) != index_identity:
raise RuntimeError("direct page pack index changed during validation")
shard_relative_paths = tuple(
(
"direct-page-packs/sha256/"
f"{_tensor_digest(shard_sha256_t).hex()}"
f"{DIRECT_PAGE_PACK_SUFFIX}"
)
for shard_sha256_t in index.shard_sha256s_t
)
shard_roots = (resolved_root,) * len(shard_relative_paths)
authority = DirectPagePackSetAuthorityPacket(
shard_roots=shard_roots,
shard_relative_paths=shard_relative_paths,
index_root=resolved_root,
index_relative_path=index_relative_path,
shard_sha256s_t=index.shard_sha256s_t.clone(),
index_sha256_t=_digest_tensor(bytes.fromhex(index_sha256)),
shard_bytes_t=index.shard_bytes_t.clone(),
index_bytes_t=torch.tensor(index_identity[2], dtype=torch.long),
logical_object_bytes_t=index.object_bytes_t.sum().reshape(()),
page_count_t=torch.tensor(
index.page_ids_t.numel(),
dtype=torch.long,
),
alignment_bytes_t=torch.tensor(
DIRECT_PAGE_PACK_ALIGNMENT_BYTES,
dtype=torch.long,
),
page_ids_sha256_t=index.page_ids_sha256_t.clone(),
page_map_sha256_t=index.page_map_sha256_t.clone(),
pack_set_sha256_t=index.pack_set_sha256_t.clone(),
)
validate_direct_page_pack_set_authority_boundary(
authority,
index=index,
)
return DirectPagePackBuildPacket(authority=authority, index=index)
def load_existing_direct_page_pack_set_for_source_boundary(
source: DirectPagePackSourcePacket,
*,
shard_roots: Sequence[Path],
index_root: Path,
) -> DirectPagePackBuildPacket | None:
"""Reopen one exact content-addressed pack without rewriting its payload.
This is a metadata and immutable-file-identity resume boundary. It
requires the existing index to name the exact source page/object map and
deterministic physical layout, then revalidates every source object's
direct schema. Callers that use this to publish generation authority must
still perform the full cold shard hash/read proof; this boundary never
substitutes cached metadata for that proof.
"""
(
page_ids_t,
object_sha256s_t,
object_bytes_t,
object_paths,
) = _validate_source_packet(source)
roots = tuple(_safe_root(Path(root), create=True) for root in shard_roots)
if not roots or len(roots) > page_ids_t.numel():
raise ValueError("direct page pack shard roots differ")
resolved_index_root = _safe_root(index_root, create=True)
index_directory = (
resolved_index_root / "direct-page-packs" / "index"
)
if not index_directory.exists():
return None
safe_index_directory = _safe_root(index_directory, create=False)
object_spans_t = (
object_bytes_t.add(DIRECT_PAGE_PACK_ALIGNMENT_BYTES - 1)
.div(DIRECT_PAGE_PACK_ALIGNMENT_BYTES, rounding_mode="floor")
.mul(DIRECT_PAGE_PACK_ALIGNMENT_BYTES)
.contiguous()
)
shard_indices_t = _balanced_contiguous_shards(
object_spans_t,
shard_count=len(roots),
)
object_offsets_t = torch.empty_like(object_spans_t)
expected_shard_bytes: list[int] = []
expected_shard_logical_bytes: list[int] = []
expected_shard_page_counts: list[int] = []
for shard_index in range(len(roots)):
rows_t = shard_indices_t.eq(shard_index).nonzero(
as_tuple=False
).reshape(-1)
spans_t = object_spans_t[rows_t]
object_offsets_t[rows_t] = torch.cat(
(
torch.zeros(1, dtype=torch.long),
torch.cumsum(spans_t[:-1], dim=0),
)
)
expected_shard_bytes.append(int(spans_t.sum()))
expected_shard_logical_bytes.append(
int(object_bytes_t[rows_t].sum())
)
expected_shard_page_counts.append(int(rows_t.numel()))
candidates: list[
tuple[Path, DirectPagePackIndexPacket]
] = []
for index_path in sorted(
safe_index_directory.glob(f"*{DIRECT_PAGE_PACK_INDEX_SUFFIX}")
):
identity = index_path.lstat()
if (
not stat.S_ISREG(identity.st_mode)
or identity.st_nlink != 1
):
continue
index_sha256 = _file_sha256(index_path)
if (
index_path.name
!= f"{index_sha256}{DIRECT_PAGE_PACK_INDEX_SUFFIX}"
):
continue
try:
existing = _load_index_file(index_path)
except RuntimeError:
continue
if (
not torch.equal(existing.page_ids_t, page_ids_t)
or not torch.equal(
existing.object_sha256s_t,
object_sha256s_t,
)
or not torch.equal(existing.object_bytes_t, object_bytes_t)
or not torch.equal(
existing.shard_indices_t,
shard_indices_t,
)
or not torch.equal(
existing.object_offsets_t,
object_offsets_t,
)
or not torch.equal(existing.object_spans_t, object_spans_t)
or existing.shard_bytes_t.tolist()
!= expected_shard_bytes
or existing.shard_logical_bytes_t.tolist()
!= expected_shard_logical_bytes
or existing.shard_page_counts_t.tolist()
!= expected_shard_page_counts
):
continue
candidates.append((index_path, existing))
if not candidates:
return None
source_revisions_t = torch.tensor(
[
_validated_existing_pack_source_revision_boundary(
path=path,
expected_page_id=int(page_id),
expected_sha256=object_sha256.numpy().tobytes(),
expected_bytes=int(object_bytes),
)
for page_id, object_sha256, object_bytes, path in zip(
page_ids_t,
object_sha256s_t,
object_bytes_t,
object_paths,
strict=True,
)
],
dtype=torch.long,
)
candidates = [
(path, index)
for path, index in candidates
if torch.equal(index.format_revisions_t, source_revisions_t)
]
if not candidates:
return None
if len(candidates) != 1:
raise RuntimeError(
"existing direct page pack source authority is ambiguous"
)
index_path, index = candidates[0]
shard_relative_paths = tuple(
(
"direct-page-packs/sha256/"
f"{_tensor_digest(index.shard_sha256s_t[shard_index]).hex()}"
f"{DIRECT_PAGE_PACK_SUFFIX}"
)
for shard_index in range(len(roots))
)
authority = DirectPagePackSetAuthorityPacket(
shard_roots=roots,
shard_relative_paths=shard_relative_paths,
index_root=resolved_index_root,
index_relative_path=str(index_path.relative_to(resolved_index_root)),
shard_sha256s_t=index.shard_sha256s_t.clone(),
index_sha256_t=_digest_tensor(
bytes.fromhex(_file_sha256(index_path))
),
shard_bytes_t=index.shard_bytes_t.clone(),
index_bytes_t=torch.tensor(
index_path.stat().st_size,
dtype=torch.long,
),
logical_object_bytes_t=index.object_bytes_t.sum().reshape(()),
page_count_t=torch.tensor(
index.page_ids_t.numel(),
dtype=torch.long,
),
alignment_bytes_t=torch.tensor(
DIRECT_PAGE_PACK_ALIGNMENT_BYTES,
dtype=torch.long,
),
page_ids_sha256_t=index.page_ids_sha256_t.clone(),
page_map_sha256_t=index.page_map_sha256_t.clone(),
pack_set_sha256_t=index.pack_set_sha256_t.clone(),
)
validate_direct_page_pack_set_authority_boundary(
authority,
index=index,
)
return DirectPagePackBuildPacket(authority=authority, index=index)
def _validated_existing_pack_source_revision_boundary(
*,
path: Path,
expected_page_id: int,
expected_sha256: bytes,
expected_bytes: int,
) -> int:
"""Validate one source identity without rereading its tensor payload."""
identity_before = _source_object_identity(
path,
expected_bytes=expected_bytes,
)
resolved = path.expanduser().resolve(strict=True)
if resolved.name != f"{expected_sha256.hex()}.safetensors":
raise RuntimeError(
"direct page pack source content-addressed name differs"
)
revision = _validate_source_object_schema(
resolved,
expected_page_id=expected_page_id,
)
if _file_identity(resolved) != identity_before:
raise RuntimeError(
"direct page pack source changed during resume validation"
)
return revision
def validate_direct_page_pack_set_authority_boundary(
authority: DirectPagePackSetAuthorityPacket,
*,
index: DirectPagePackIndexPacket | None = None,
) -> DirectPagePackIndexPacket:
"""Validate path containment, immutable identities, and tensor authority."""
shard_count = len(authority.shard_roots)
if (
shard_count < 1
or len(authority.shard_relative_paths) != shard_count
or authority.shard_sha256s_t.dtype != torch.uint8
or authority.shard_sha256s_t.shape != (shard_count, 32)
or authority.shard_bytes_t.dtype != torch.long
or authority.shard_bytes_t.shape != (shard_count,)
or authority.index_sha256_t.dtype != torch.uint8
or authority.index_sha256_t.reshape(-1).shape != (32,)
or authority.index_bytes_t.dtype != torch.long
or authority.index_bytes_t.reshape(-1).shape != (1,)
or int(authority.index_bytes_t) < 1
or authority.logical_object_bytes_t.dtype != torch.long
or authority.logical_object_bytes_t.reshape(-1).shape != (1,)
or int(authority.logical_object_bytes_t) < 1
or authority.page_count_t.dtype != torch.long
or authority.page_count_t.reshape(-1).shape != (1,)
or int(authority.page_count_t) < 1
or authority.alignment_bytes_t.dtype != torch.long
or authority.alignment_bytes_t.reshape(-1).shape != (1,)
or int(authority.alignment_bytes_t)
!= DIRECT_PAGE_PACK_ALIGNMENT_BYTES
):
raise RuntimeError("direct page pack set authority tensor geometry differs")
index_path = _safe_relative_path(
authority.index_root,
authority.index_relative_path,
must_exist=True,
)
index_identity = _file_identity(index_path)
if (
index_identity[2] != int(authority.index_bytes_t)
or bytes.fromhex(_file_sha256(index_path))
!= _tensor_digest(authority.index_sha256_t)
):
raise RuntimeError("direct page pack index file authority differs")
loaded_index = _load_index_file(index_path)
if index is not None:
supplied_index = validate_direct_page_pack_index_boundary(index)
for field_name in DirectPagePackIndexPacket.__dataclass_fields__:
loaded_value = getattr(loaded_index, field_name)
supplied_value = getattr(supplied_index, field_name)
if not torch.equal(loaded_value, supplied_value):
raise RuntimeError(
"direct page pack supplied index differs from its file"
)
if _file_identity(index_path) != index_identity:
raise RuntimeError("direct page pack index changed during validation")
if (
not torch.equal(
loaded_index.shard_sha256s_t,
authority.shard_sha256s_t.detach().cpu(),
)
or not torch.equal(
loaded_index.shard_bytes_t,
authority.shard_bytes_t.detach().cpu(),
)
or int(loaded_index.object_bytes_t.sum())
!= int(authority.logical_object_bytes_t)
or loaded_index.page_ids_t.numel() != int(authority.page_count_t)
or not torch.equal(
loaded_index.page_ids_sha256_t,
authority.page_ids_sha256_t.detach().cpu(),
)
or not torch.equal(
loaded_index.page_map_sha256_t,
authority.page_map_sha256_t.detach().cpu(),
)
or not torch.equal(
loaded_index.pack_set_sha256_t,
authority.pack_set_sha256_t.detach().cpu(),
)
):
raise RuntimeError("direct page pack set authority differs from its index")
for shard_index, (root, relative_path) in enumerate(
zip(
authority.shard_roots,
authority.shard_relative_paths,
strict=True,
)
):
shard_path = _safe_relative_path(
root,
relative_path,
must_exist=True,
)
identity = shard_path.lstat()
if (
identity.st_size != int(loaded_index.shard_bytes_t[shard_index])
or identity.st_size % DIRECT_PAGE_PACK_ALIGNMENT_BYTES
or identity.st_nlink != 1
):
raise RuntimeError("direct page pack shard file authority differs")
return loaded_index
def load_direct_page_pack_index_boundary(
authority: DirectPagePackSetAuthorityPacket,
) -> DirectPagePackIndexPacket:
return validate_direct_page_pack_set_authority_boundary(authority)
def _open_pack_descriptor(path: Path, *, writable: bool, direct_io: bool) -> int:
flags = (
(os.O_WRONLY if writable else os.O_RDONLY)
| getattr(os, "O_CLOEXEC", 0)
| getattr(os, "O_NOFOLLOW", 0)
)
if direct_io:
direct_flag = getattr(os, "O_DIRECT", 0)
if not direct_flag:
raise RuntimeError("direct page pack O_DIRECT is unavailable")
flags |= direct_flag
if writable:
flags |= os.O_CREAT | os.O_EXCL
try:
return os.open(path, flags, 0o600)
except OSError as error:
raise RuntimeError("direct page pack descriptor open failed") from error
def _pread_exact_into(
descriptor: int,
view: memoryview,
*,
offset: int,
direct_io: bool,
) -> None:
completed = 0
while completed < len(view):
tail = view[completed:]
try:
count = os.preadv(descriptor, (tail,), offset + completed)
finally:
tail.release()
if count < 1:
raise RuntimeError("direct page pack read was incomplete")
completed += count
if (
direct_io
and completed < len(view)
and completed % DIRECT_PAGE_PACK_ALIGNMENT_BYTES
):
raise RuntimeError("direct page pack direct read lost alignment")
def _pwrite_exact_from(
descriptor: int,
view: memoryview,
*,
offset: int,
) -> None:
completed = 0
while completed < len(view):
tail = view[completed:]
try:
count = os.pwritev(descriptor, (tail,), offset + completed)
finally:
tail.release()
if count < 1:
raise RuntimeError("direct page pack write was incomplete")
completed += count
if (
completed < len(view)
and completed % DIRECT_PAGE_PACK_ALIGNMENT_BYTES
):
raise RuntimeError("direct page pack direct write lost alignment")
def _shard_path(
authority: DirectPagePackSetAuthorityPacket,
shard_index: int,
) -> Path:
return _safe_relative_path(
authority.shard_roots[shard_index],
authority.shard_relative_paths[shard_index],
must_exist=True,
)
def _read_selected_shard(
*,
authority: DirectPagePackSetAuthorityPacket,
index: DirectPagePackIndexPacket,
shard_index: int,
requested_rows: tuple[tuple[int, int], ...],
direct_io: bool,
) -> tuple[tuple[tuple[int, bytes], ...], int]:
shard_path = _shard_path(authority, shard_index)
identity_before = _file_identity(shard_path)
descriptor = _open_pack_descriptor(
shard_path,
writable=False,
direct_io=direct_io,
)
ordered_rows = tuple(
sorted(
requested_rows,
key=lambda pair: int(index.object_offsets_t[pair[1]]),
)
)
groups: list[list[tuple[int, int]]] = []
for request in ordered_rows:
row_index = request[1]
row_span = int(index.object_spans_t[row_index])
if (
groups
and int(index.object_offsets_t[row_index])
== int(index.object_offsets_t[groups[-1][-1][1]])
+ int(index.object_spans_t[groups[-1][-1][1]])
and (
int(index.object_offsets_t[row_index])
+ row_span
- int(index.object_offsets_t[groups[-1][0][1]])
<= DIRECT_PAGE_PACK_MAX_COALESCED_READ_BYTES
)
):
groups[-1].append(request)
else:
groups.append([request])
payloads: list[tuple[int, bytes]] = []
physical_read_bytes = 0
try:
for group in groups:
first_row = group[0][1]
last_row = group[-1][1]
group_offset = int(index.object_offsets_t[first_row])
group_end = (
int(index.object_offsets_t[last_row])
+ int(index.object_spans_t[last_row])
)
group_bytes = group_end - group_offset
bounce = mmap.mmap(-1, group_bytes, access=mmap.ACCESS_WRITE)
view = memoryview(bounce)
try:
_pread_exact_into(
descriptor,
view,
offset=group_offset,
direct_io=direct_io,
)
physical_read_bytes += group_bytes
for request_index, row_index in group:
local_offset = (
int(index.object_offsets_t[row_index]) - group_offset
)
logical_bytes = int(index.object_bytes_t[row_index])
span = int(index.object_spans_t[row_index])
payload_view = view[
local_offset : local_offset + logical_bytes
]
padding_view = view[
local_offset + logical_bytes : local_offset + span
]
try:
expected_sha256 = (
index.object_sha256s_t[row_index]
.numpy()
.tobytes(order="C")
)
if hashlib.sha256(payload_view).digest() != expected_sha256:
raise RuntimeError(
"direct page pack selected object hash differs"
)
if padding_view and padding_view.tobytes().strip(b"\x00"):
raise RuntimeError(
"direct page pack selected padding is nonzero"
)
payloads.append(
(request_index, payload_view.tobytes())
)
finally:
payload_view.release()
padding_view.release()
finally:
view.release()
bounce.close()
finally:
os.close(descriptor)
if _file_identity(shard_path) != identity_before:
raise RuntimeError("direct page pack shard changed during selected read")
return tuple(payloads), physical_read_bytes
def read_direct_page_pack_selected_boundary(
authority: DirectPagePackSetAuthorityPacket,
selected_page_ids_t: torch.Tensor,
*,
direct_io: bool = True,
) -> DirectPagePackSelectedPacket:
"""Read selected complete objects, coalescing within and parallelizing shards."""
index = load_direct_page_pack_index_boundary(authority)
selected_t = _stable_cpu_tensor(
selected_page_ids_t,
dtype=torch.long,
).reshape(-1)
if (
selected_page_ids_t.dtype != torch.long
or selected_t.numel() < 1
or torch.unique(selected_t).numel() != selected_t.numel()
):
raise RuntimeError("direct page pack selection differs")
positions_t = torch.searchsorted(index.page_ids_t, selected_t)
bounded_positions_t = positions_t.clamp(max=index.page_ids_t.numel() - 1)
if bool(positions_t.ge(index.page_ids_t.numel()).any()) or not torch.equal(
index.page_ids_t[bounded_positions_t],
selected_t,
):
raise FileNotFoundError("selected direct page is absent from the pack set")
requested_by_shard: dict[int, list[tuple[int, int]]] = {}
for request_index, row_index_t in enumerate(positions_t):
row_index = int(row_index_t)
shard_index = int(index.shard_indices_t[row_index])
requested_by_shard.setdefault(shard_index, []).append(
(request_index, row_index)
)
shard_physical_read_bytes_t = torch.zeros(
index.shard_sha256s_t.shape[0],
dtype=torch.long,
)
payload_by_request: dict[int, bytes] = {}
with ThreadPoolExecutor(
max_workers=len(requested_by_shard),
thread_name_prefix="nnf-direct-pack-read",
) as executor:
futures = {
shard_index: executor.submit(
_read_selected_shard,
authority=authority,
index=index,
shard_index=shard_index,
requested_rows=tuple(requested_rows),
direct_io=direct_io,
)
for shard_index, requested_rows in requested_by_shard.items()
}
for shard_index, future in futures.items():
payloads, physical_read_bytes = future.result()
shard_physical_read_bytes_t[shard_index] = physical_read_bytes
payload_by_request.update(payloads)
selected_bytes_t = index.object_bytes_t[positions_t].clone()
payload_offsets_t = torch.cat(
(
torch.zeros(1, dtype=torch.long),
torch.cumsum(selected_bytes_t, dim=0),
)
)
object_payload_t = torch.empty(
int(payload_offsets_t[-1]),
dtype=torch.uint8,
)
for request_index in range(selected_t.numel()):
payload = payload_by_request.get(request_index)
if payload is None:
raise RuntimeError("direct page pack selected payload is absent")
target_t = object_payload_t[
int(payload_offsets_t[request_index]) : int(
payload_offsets_t[request_index + 1]
)
]
target_t.copy_(
torch.frombuffer(bytearray(payload), dtype=torch.uint8)
)
return DirectPagePackSelectedPacket(
page_ids_t=selected_t,
object_sha256s_t=index.object_sha256s_t[positions_t].clone(),
object_bytes_t=selected_bytes_t,
payload_offsets_t=payload_offsets_t,
object_payload_t=object_payload_t,
shard_physical_read_bytes_t=shard_physical_read_bytes_t,
direct_io_t=torch.tensor(direct_io, dtype=torch.bool),
)
def _verify_one_shard_cold(
*,
authority: DirectPagePackSetAuthorityPacket,
index: DirectPagePackIndexPacket,
shard_index: int,
direct_io: bool,
) -> tuple[int, int]:
shard_path = _shard_path(authority, shard_index)
identity_before = _file_identity(shard_path)
shard_bytes = int(index.shard_bytes_t[shard_index])
descriptor = _open_pack_descriptor(
shard_path,
writable=False,
direct_io=direct_io,
)
bounce_bytes = min(DIRECT_PAGE_PACK_IO_WAVE_BYTES, shard_bytes)
bounce = mmap.mmap(-1, bounce_bytes, access=mmap.ACCESS_WRITE)
view = memoryview(bounce)
shard_digest = hashlib.sha256()
row_indices = tuple(
int(row_index)
for row_index in index.shard_indices_t.eq(shard_index)
.nonzero(as_tuple=False)
.reshape(-1)
)
row_position = 0
object_digest = hashlib.sha256()
object_verified = False
offset = 0
started_ns = time.perf_counter_ns()
try:
while offset < shard_bytes:
wave_bytes = min(bounce_bytes, shard_bytes - offset)
wave = view[:wave_bytes]
try:
_pread_exact_into(
descriptor,
wave,
offset=offset,
direct_io=direct_io,
)
shard_digest.update(wave)
wave_end = offset + wave_bytes
cursor = offset
while cursor < wave_end:
if row_position >= len(row_indices):
raise RuntimeError(
"direct page pack cold object layout differs"
)
row_index = row_indices[row_position]
object_start = int(index.object_offsets_t[row_index])
object_end = object_start + int(
index.object_bytes_t[row_index]
)
span_end = object_start + int(
index.object_spans_t[row_index]
)
if cursor < object_start or cursor >= span_end:
raise RuntimeError(
"direct page pack cold object layout differs"
)
if cursor < object_end:
payload_end = min(wave_end, object_end)
payload_view = wave[
cursor - offset : payload_end - offset
]
try:
object_digest.update(payload_view)
finally:
payload_view.release()
cursor = payload_end
if cursor < object_end:
continue
if not object_verified:
if object_digest.digest() != _tensor_digest(
index.object_sha256s_t[row_index]
):
raise RuntimeError(
"direct page pack cold object hash differs"
)
object_verified = True
if cursor < span_end:
padding_end = min(wave_end, span_end)
padding_view = wave[
cursor - offset : padding_end - offset
]
try:
if padding_view.tobytes().strip(b"\x00"):
raise RuntimeError(
"direct page pack cold padding is nonzero"
)
finally:
padding_view.release()
cursor = padding_end
if cursor < span_end:
continue
row_position += 1
object_digest = hashlib.sha256()
object_verified = False
finally:
wave.release()
offset += wave_bytes
finally:
view.release()
bounce.close()
os.close(descriptor)
elapsed_ns = time.perf_counter_ns() - started_ns
if row_position != len(row_indices):
raise RuntimeError("direct page pack cold object layout differs")
if shard_digest.digest() != _tensor_digest(
index.shard_sha256s_t[shard_index]
):
raise RuntimeError("direct page pack cold shard hash differs")
if _file_identity(shard_path) != identity_before:
raise RuntimeError("direct page pack shard changed during cold read")
return shard_bytes, elapsed_ns
def verify_direct_page_pack_set_cold_boundary(
authority: DirectPagePackSetAuthorityPacket,
*,
minimum_unique_logical_bytes_per_second: int = 0,
direct_io: bool = True,
) -> DirectPagePackColdVerificationPacket:
"""Cold-read every unique shard once and report parallel ingestion speed."""
if minimum_unique_logical_bytes_per_second < 0:
raise ValueError("direct page pack cold minimum rate differs")
index = load_direct_page_pack_index_boundary(authority)
shard_count = index.shard_sha256s_t.shape[0]
wall_started_ns = time.perf_counter_ns()
with ThreadPoolExecutor(
max_workers=shard_count,
thread_name_prefix="nnf-direct-pack-cold",
) as executor:
futures = tuple(
executor.submit(
_verify_one_shard_cold,
authority=authority,
index=index,
shard_index=shard_index,
direct_io=direct_io,
)
for shard_index in range(shard_count)
)
rows = tuple(future.result() for future in futures)
wall_elapsed_ns = time.perf_counter_ns() - wall_started_ns
shard_physical_read_bytes_t = torch.tensor(
[row[0] for row in rows],
dtype=torch.long,
)
shard_logical_rates_t = torch.tensor(
[
int(index.shard_logical_bytes_t[shard_index])
* 1_000_000_000
// max(1, rows[shard_index][1])
for shard_index in range(shard_count)
],
dtype=torch.long,
)
logical_object_bytes = int(index.shard_logical_bytes_t.sum())
aggregate_rate = (
logical_object_bytes * 1_000_000_000 // max(1, wall_elapsed_ns)
)
if aggregate_rate < minimum_unique_logical_bytes_per_second:
raise RuntimeError(
"direct page pack unique cold-ingestion rate is below authority"
)
return DirectPagePackColdVerificationPacket(
pack_set_sha256_t=index.pack_set_sha256_t.clone(),
shard_sha256s_t=index.shard_sha256s_t.clone(),
shard_physical_read_bytes_t=shard_physical_read_bytes_t,
shard_logical_bytes_per_second_t=shard_logical_rates_t,
logical_object_bytes_t=torch.tensor(
logical_object_bytes,
dtype=torch.long,
),
physical_read_bytes_t=shard_physical_read_bytes_t.sum().reshape(()),
elapsed_nanoseconds_t=torch.tensor(wall_elapsed_ns, dtype=torch.long),
aggregate_unique_logical_bytes_per_second_t=torch.tensor(
aggregate_rate,
dtype=torch.long,
),
direct_io_t=torch.tensor(direct_io, dtype=torch.bool),
zero_padding_verified_t=torch.tensor(True, dtype=torch.bool),
)
def _copy_file_buffered_atomic(
source: Path,
destination: Path,
*,
expected_sha256: bytes,
) -> bool:
if destination.exists():
if (
destination.is_symlink()
or bytes.fromhex(_file_sha256(destination)) != expected_sha256
):
raise RuntimeError("existing direct page pack index replica differs")
return False
temporary = destination.parent / (
f".{destination.name}.{os.getpid()}.{time.monotonic_ns()}.tmp"
)
descriptor = os.open(
temporary,
os.O_WRONLY
| os.O_CREAT
| os.O_EXCL
| getattr(os, "O_CLOEXEC", 0)
| getattr(os, "O_NOFOLLOW", 0),
0o600,
)
digest = hashlib.sha256()
try:
with source.open("rb") as source_handle:
while True:
chunk = source_handle.read(1024 * 1024)
if not chunk:
break
digest.update(chunk)
written = 0
while written < len(chunk):
count = os.write(descriptor, chunk[written:])
if count < 1:
raise RuntimeError(
"direct page pack index replica write was incomplete"
)
written += count
os.fsync(descriptor)
finally:
os.close(descriptor)
try:
if digest.digest() != expected_sha256:
raise RuntimeError("direct page pack index source hash differs")
if destination.exists():
raise RuntimeError(
"direct page pack index replica appeared during publication"
)
os.rename(temporary, destination)
_fsync_directory(destination.parent)
finally:
temporary.unlink(missing_ok=True)
return True
def _copy_one_shard_direct(
*,
source_path: Path,
destination_path: Path,
expected_sha256: bytes,
physical_bytes: int,
logical_bytes: int,
) -> tuple[bool, int, int]:
if destination_path.exists():
identity_before = _file_identity(destination_path)
if (
destination_path.is_symlink()
or identity_before[2] != physical_bytes
):
raise RuntimeError(
"existing direct page pack shard replica differs"
)
descriptor = _open_pack_descriptor(
destination_path,
writable=False,
direct_io=True,
)
bounce = mmap.mmap(
-1,
min(DIRECT_PAGE_PACK_IO_WAVE_BYTES, physical_bytes),
access=mmap.ACCESS_WRITE,
)
view = memoryview(bounce)
digest = hashlib.sha256()
offset = 0
try:
while offset < physical_bytes:
wave_bytes = min(len(view), physical_bytes - offset)
wave = view[:wave_bytes]
try:
_pread_exact_into(
descriptor,
wave,
offset=offset,
direct_io=True,
)
digest.update(wave)
finally:
wave.release()
offset += wave_bytes
finally:
view.release()
bounce.close()
os.close(descriptor)
if (
digest.digest() != expected_sha256
or _file_identity(destination_path) != identity_before
):
raise RuntimeError(
"existing direct page pack shard replica differs"
)
return False, 0, 0
temporary = destination_path.parent / (
f".{destination_path.name}.{os.getpid()}.{time.monotonic_ns()}.tmp"
)
source_descriptor = -1
destination_descriptor = -1
try:
source_descriptor = _open_pack_descriptor(
source_path,
writable=False,
direct_io=True,
)
destination_descriptor = _open_pack_descriptor(
temporary,
writable=True,
direct_io=True,
)
if hasattr(os, "posix_fallocate"):
os.posix_fallocate(destination_descriptor, 0, physical_bytes)
else:
os.ftruncate(destination_descriptor, physical_bytes)
bounce = mmap.mmap(
-1,
min(DIRECT_PAGE_PACK_IO_WAVE_BYTES, physical_bytes),
access=mmap.ACCESS_WRITE,
)
view = memoryview(bounce)
digest = hashlib.sha256()
offset = 0
started_ns = time.perf_counter_ns()
try:
while offset < physical_bytes:
wave_bytes = min(len(view), physical_bytes - offset)
wave = view[:wave_bytes]
try:
_pread_exact_into(
source_descriptor,
wave,
offset=offset,
direct_io=True,
)
digest.update(wave)
_pwrite_exact_from(
destination_descriptor,
wave,
offset=offset,
)
finally:
wave.release()
offset += wave_bytes
os.fsync(destination_descriptor)
finally:
view.release()
bounce.close()
elapsed_ns = time.perf_counter_ns() - started_ns
os.close(destination_descriptor)
destination_descriptor = -1
os.close(source_descriptor)
source_descriptor = -1
if digest.digest() != expected_sha256:
raise RuntimeError("direct page pack replica source hash differs")
if destination_path.exists():
raise RuntimeError(
"direct page pack shard replica appeared during publication"
)
os.rename(temporary, destination_path)
_fsync_directory(destination_path.parent)
return True, logical_bytes, elapsed_ns
finally:
if destination_descriptor >= 0:
os.close(destination_descriptor)
if source_descriptor >= 0:
os.close(source_descriptor)
temporary.unlink(missing_ok=True)
def replicate_direct_page_pack_set_durable_boundary(
authority: DirectPagePackSetAuthorityPacket,
*,
destination_shard_roots: Sequence[Path],
destination_index_root: Path,
minimum_unique_logical_bytes_per_second: int = 0,
) -> DirectPagePackReplicaReceiptPacket:
"""Durably copy each unique shard once through O_DIRECT.
This API creates one destination pack set. It records durable-copy speed
but does not apply the 1 GB/s cold-ingestion gate by default. It never
multiplies the numerator by replica count: each disjoint logical shard
contributes at most its own logical object bytes.
"""
if minimum_unique_logical_bytes_per_second < 0:
raise ValueError("direct page pack replica minimum rate differs")
index = load_direct_page_pack_index_boundary(authority)
destination_roots = tuple(
_safe_root(Path(root), create=True) for root in destination_shard_roots
)
if len(destination_roots) != len(authority.shard_roots):
raise ValueError("direct page pack destination shard count differs")
destination_paths = tuple(
_safe_relative_path(
destination_roots[shard_index],
authority.shard_relative_paths[shard_index],
must_exist=False,
create_parent=True,
)
for shard_index in range(len(destination_roots))
)
created_paths: list[Path] = []
index_created_path: Path | None = None
wall_started_ns = time.perf_counter_ns()
try:
with ThreadPoolExecutor(
max_workers=len(destination_roots),
thread_name_prefix="nnf-direct-pack-replica",
) as executor:
futures = tuple(
executor.submit(
_copy_one_shard_direct,
source_path=_shard_path(authority, shard_index),
destination_path=destination_paths[shard_index],
expected_sha256=_tensor_digest(
index.shard_sha256s_t[shard_index]
),
physical_bytes=int(index.shard_bytes_t[shard_index]),
logical_bytes=int(
index.shard_logical_bytes_t[shard_index]
),
)
for shard_index in range(len(destination_roots))
)
shard_rows = tuple(future.result() for future in futures)
for shard_index, row in enumerate(shard_rows):
if row[0]:
created_paths.append(destination_paths[shard_index])
resolved_destination_index_root = _safe_root(
destination_index_root,
create=True,
)
destination_index_path = _safe_relative_path(
resolved_destination_index_root,
authority.index_relative_path,
must_exist=False,
create_parent=True,
)
source_index_path = _safe_relative_path(
authority.index_root,
authority.index_relative_path,
must_exist=True,
)
if _copy_file_buffered_atomic(
source_index_path,
destination_index_path,
expected_sha256=_tensor_digest(authority.index_sha256_t),
):
index_created_path = destination_index_path
wall_elapsed_ns = time.perf_counter_ns() - wall_started_ns
newly_written_logical_bytes_t = torch.tensor(
[row[1] for row in shard_rows],
dtype=torch.long,
)
newly_written_shards_t = newly_written_logical_bytes_t.gt(0)
total_newly_written = int(newly_written_logical_bytes_t.sum())
aggregate_rate = (
total_newly_written * 1_000_000_000 // max(1, wall_elapsed_ns)
)
if (
total_newly_written < 1
and minimum_unique_logical_bytes_per_second > 0
) or aggregate_rate < minimum_unique_logical_bytes_per_second:
raise RuntimeError(
"direct page pack durable replica rate is below authority"
)
shard_rates_t = torch.tensor(
[
row[1] * 1_000_000_000 // max(1, row[2])
if row[1] > 0
else 0
for row in shard_rows
],
dtype=torch.long,
)
destination_authority = replace(
authority,
shard_roots=destination_roots,
index_root=resolved_destination_index_root,
)
validate_direct_page_pack_set_authority_boundary(
destination_authority
)
cold_verification = verify_direct_page_pack_set_cold_boundary(
destination_authority,
minimum_unique_logical_bytes_per_second=0,
direct_io=True,
)
return DirectPagePackReplicaReceiptPacket(
authority=destination_authority,
newly_written_shards_t=newly_written_shards_t,
newly_written_logical_bytes_t=newly_written_logical_bytes_t,
shard_replica_bytes_per_second_t=shard_rates_t,
total_newly_written_logical_bytes_t=torch.tensor(
total_newly_written,
dtype=torch.long,
),
elapsed_nanoseconds_t=torch.tensor(
wall_elapsed_ns,
dtype=torch.long,
),
aggregate_unique_logical_bytes_per_second_t=torch.tensor(
aggregate_rate,
dtype=torch.long,
),
direct_io_t=torch.tensor(True, dtype=torch.bool),
cold_verification=cold_verification,
)
except Exception:
touched_roots: set[Path] = set()
for path in reversed(created_paths):
if path.exists() and not path.is_symlink():
path.unlink()
touched_roots.add(path.parent)
if (
index_created_path is not None
and index_created_path.exists()
and not index_created_path.is_symlink()
):
index_created_path.unlink()
touched_roots.add(index_created_path.parent)
for root in sorted(touched_roots, key=str):
_fsync_directory(root)
raise
def replicate_direct_page_pack_shard_durable_boundary(
authority: DirectPagePackSetAuthorityPacket,
*,
shard_index: int,
destination_root: Path,
) -> DirectPagePackShardReplicaReceiptPacket:
"""Durably copy or verify one proof-only shard through O_DIRECT."""
index = load_direct_page_pack_index_boundary(authority)
if (
isinstance(shard_index, bool)
or shard_index < 0
or shard_index >= index.shard_sha256s_t.shape[0]
):
raise ValueError("direct page pack proof shard index differs")
root = _safe_root(destination_root, create=True)
destination_path = _safe_relative_path(
root,
authority.shard_relative_paths[shard_index],
must_exist=False,
create_parent=True,
)
written, logical_bytes, elapsed_ns = _copy_one_shard_direct(
source_path=_shard_path(authority, shard_index),
destination_path=destination_path,
expected_sha256=_tensor_digest(
index.shard_sha256s_t[shard_index]
),
physical_bytes=int(index.shard_bytes_t[shard_index]),
logical_bytes=int(index.shard_logical_bytes_t[shard_index]),
)
return DirectPagePackShardReplicaReceiptPacket(
shard_index_t=torch.tensor(shard_index, dtype=torch.long),
shard_root=root,
shard_relative_path=authority.shard_relative_paths[shard_index],
shard_sha256_t=index.shard_sha256s_t[shard_index].clone(),
shard_bytes_t=index.shard_bytes_t[shard_index].clone(),
logical_object_bytes_t=index.shard_logical_bytes_t[
shard_index
].clone(),
newly_written_t=torch.tensor(written, dtype=torch.bool),
newly_written_logical_bytes_t=torch.tensor(
logical_bytes,
dtype=torch.long,
),
elapsed_nanoseconds_t=torch.tensor(
elapsed_ns,
dtype=torch.long,
),
direct_io_t=torch.tensor(True, dtype=torch.bool),
)