feat(simulation): upload real gaussian source bundles
This commit is contained in:
@@ -46,23 +46,39 @@ class GaussianPipelineIntegrityError(GaussianPipelineGatewayError):
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GaussianSourceUpload:
|
||||
class GaussianSourceMemberUpload:
|
||||
upload_id: str
|
||||
filename: str
|
||||
format: str
|
||||
logical_path: str
|
||||
sha256: str
|
||||
byte_length: int
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
"upload_id": self.upload_id,
|
||||
"filename": self.filename,
|
||||
"format": self.format,
|
||||
"logical_path": self.logical_path,
|
||||
"sha256": self.sha256,
|
||||
"byte_length": self.byte_length,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GaussianSourceBundleUpload:
|
||||
format: str
|
||||
entrypoint: str
|
||||
bundle_sha256: str
|
||||
total_byte_length: int
|
||||
members: tuple[GaussianSourceMemberUpload, ...]
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
"format": self.format,
|
||||
"entrypoint": self.entrypoint,
|
||||
"bundle_sha256": self.bundle_sha256,
|
||||
"total_byte_length": self.total_byte_length,
|
||||
"members": [member.to_dict() for member in self.members],
|
||||
}
|
||||
|
||||
|
||||
class GaussianPipelineGateway:
|
||||
"""TUS and JSON client with bounded responses and independent digest checks."""
|
||||
|
||||
@@ -106,37 +122,90 @@ class GaussianPipelineGateway:
|
||||
document.get("schema_version") != CAPABILITIES_SCHEMA
|
||||
or document.get("service") != "ndc-gaussian-pipeline"
|
||||
or document.get("api_version") != "gaussian-pipeline.api/v1"
|
||||
or document.get("upload_protocol") != "tus/1.0.0"
|
||||
or document.get("source_transport") != "tus-bundle/v1"
|
||||
):
|
||||
raise GaussianPipelineGatewayError("Gaussian provider capabilities do not match v1")
|
||||
_validate_runtime_provenance(document)
|
||||
return document
|
||||
|
||||
def upload_source(
|
||||
def upload_source_bundle(
|
||||
self,
|
||||
source_path: Path,
|
||||
bundle_root: Path,
|
||||
*,
|
||||
filename: str,
|
||||
entrypoint: str,
|
||||
source_format: str,
|
||||
sha256: str,
|
||||
) -> GaussianSourceUpload:
|
||||
source_candidate = source_path.expanduser().absolute()
|
||||
if source_candidate.is_symlink() or not source_candidate.is_file():
|
||||
raise GaussianPipelineIntegrityError("Gaussian source must be one regular file")
|
||||
source = source_candidate.resolve()
|
||||
) -> GaussianSourceBundleUpload:
|
||||
if source_format not in {"lcc", "lcc2"}:
|
||||
raise GaussianPipelineIntegrityError("Gaussian source format must be lcc or lcc2")
|
||||
if Path(filename).name != filename or not filename.lower().endswith(f".{source_format}"):
|
||||
raise GaussianPipelineIntegrityError("Gaussian source filename is unsafe")
|
||||
if SHA256_PATTERN.fullmatch(sha256) is None:
|
||||
raise GaussianPipelineIntegrityError("Gaussian source digest is invalid")
|
||||
byte_length = source.stat().st_size
|
||||
if byte_length <= 0:
|
||||
raise GaussianPipelineIntegrityError("Gaussian source is empty")
|
||||
if _sha256(source) != sha256:
|
||||
raise GaussianPipelineIntegrityError("Gaussian source digest does not match")
|
||||
logical_entrypoint = _logical_path(entrypoint, "entrypoint")
|
||||
if not logical_entrypoint.lower().endswith(f".{source_format}"):
|
||||
raise GaussianPipelineIntegrityError(
|
||||
"Gaussian source entrypoint extension does not match its format"
|
||||
)
|
||||
root_candidate = bundle_root.expanduser().absolute()
|
||||
if root_candidate.is_symlink() or not root_candidate.is_dir():
|
||||
raise GaussianPipelineIntegrityError(
|
||||
"Gaussian source bundle root must be one regular directory"
|
||||
)
|
||||
root = root_candidate.resolve()
|
||||
logical_paths = _discover_bundle_members(root, logical_entrypoint, source_format)
|
||||
local_members: list[tuple[str, Path, str, int]] = []
|
||||
total_byte_length = 0
|
||||
for logical_path in sorted(logical_paths, key=lambda value: value.encode("utf-8")):
|
||||
source = _bundle_member(root, logical_path)
|
||||
byte_length = source.stat().st_size
|
||||
if byte_length <= 0:
|
||||
raise GaussianPipelineIntegrityError(
|
||||
f"Gaussian source bundle member is empty: {logical_path}"
|
||||
)
|
||||
digest = _sha256(source)
|
||||
total_byte_length += byte_length
|
||||
local_members.append((logical_path, source, digest, byte_length))
|
||||
|
||||
capabilities = self.capabilities()
|
||||
max_source_bytes = capabilities.get("max_source_bytes")
|
||||
max_source_files = capabilities.get("max_source_files")
|
||||
if (
|
||||
not isinstance(max_source_bytes, int)
|
||||
or isinstance(max_source_bytes, bool)
|
||||
or total_byte_length > max_source_bytes
|
||||
):
|
||||
raise GaussianPipelineIntegrityError(
|
||||
"Gaussian source bundle exceeds provider byte admission"
|
||||
)
|
||||
if (
|
||||
not isinstance(max_source_files, int)
|
||||
or isinstance(max_source_files, bool)
|
||||
or len(local_members) > max_source_files
|
||||
):
|
||||
raise GaussianPipelineIntegrityError(
|
||||
"Gaussian source bundle exceeds provider file admission"
|
||||
)
|
||||
|
||||
members = tuple(
|
||||
self._upload_member(source, logical_path, digest, byte_length)
|
||||
for logical_path, source, digest, byte_length in local_members
|
||||
)
|
||||
bundle_sha256 = _bundle_sha256(source_format, logical_entrypoint, members)
|
||||
return GaussianSourceBundleUpload(
|
||||
format=source_format,
|
||||
entrypoint=logical_entrypoint,
|
||||
bundle_sha256=bundle_sha256,
|
||||
total_byte_length=total_byte_length,
|
||||
members=members,
|
||||
)
|
||||
|
||||
def _upload_member(
|
||||
self,
|
||||
source: Path,
|
||||
logical_path: str,
|
||||
sha256: str,
|
||||
byte_length: int,
|
||||
) -> GaussianSourceMemberUpload:
|
||||
|
||||
metadata = _tus_metadata(
|
||||
{"filename": filename, "format": source_format, "sha256": sha256}
|
||||
{"logical_path": logical_path, "sha256": sha256}
|
||||
)
|
||||
try:
|
||||
response = self._client.post(
|
||||
@@ -161,10 +230,9 @@ class GaussianPipelineGateway:
|
||||
if SAFE_UPLOAD_ID.fullmatch(upload_id) is None:
|
||||
raise GaussianPipelineGatewayError("Gaussian upload id is invalid")
|
||||
self._send_file(source, upload_url, byte_length)
|
||||
return GaussianSourceUpload(
|
||||
return GaussianSourceMemberUpload(
|
||||
upload_id=upload_id,
|
||||
filename=filename,
|
||||
format=source_format,
|
||||
logical_path=logical_path,
|
||||
sha256=sha256,
|
||||
byte_length=byte_length,
|
||||
)
|
||||
@@ -399,6 +467,134 @@ def _sha256(path: Path) -> str:
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _bundle_sha256(
|
||||
source_format: str,
|
||||
entrypoint: str,
|
||||
members: tuple[GaussianSourceMemberUpload, ...],
|
||||
) -> str:
|
||||
document = {
|
||||
"format": source_format,
|
||||
"entrypoint": entrypoint,
|
||||
"members": [
|
||||
{
|
||||
"logical_path": member.logical_path,
|
||||
"sha256": member.sha256,
|
||||
"byte_length": member.byte_length,
|
||||
}
|
||||
for member in members
|
||||
],
|
||||
}
|
||||
canonical = json.dumps(
|
||||
document,
|
||||
ensure_ascii=False,
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
return hashlib.sha256(canonical).hexdigest()
|
||||
|
||||
|
||||
def _discover_bundle_members(root: Path, entrypoint: str, source_format: str) -> set[str]:
|
||||
descriptor = _bundle_member(root, entrypoint)
|
||||
if descriptor.stat().st_size > 16 * 1024 * 1024:
|
||||
raise GaussianPipelineIntegrityError("Gaussian source descriptor is too large")
|
||||
try:
|
||||
text = descriptor.read_text(encoding="utf-8")
|
||||
try:
|
||||
document = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
document = json.loads(re.sub(r",(?=\s*[}\]])", "", text))
|
||||
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||||
raise GaussianPipelineIntegrityError(
|
||||
"Gaussian source descriptor is not readable JSON"
|
||||
) from exc
|
||||
if not isinstance(document, dict):
|
||||
raise GaussianPipelineIntegrityError("Gaussian source descriptor must be an object")
|
||||
|
||||
base = entrypoint.rsplit("/", 1)[0] if "/" in entrypoint else ""
|
||||
|
||||
def related(name: str) -> str:
|
||||
logical = f"{base}/{name}" if base else name
|
||||
return _logical_path(logical, "descriptor member")
|
||||
|
||||
members = {entrypoint}
|
||||
if source_format == "lcc":
|
||||
members.update({related("index.bin"), related("data.bin")})
|
||||
file_type = document.get("fileType")
|
||||
attributes = document.get("attributes")
|
||||
has_sh = file_type == "Quality"
|
||||
if file_type not in {"Portable", "Quality"}:
|
||||
has_sh = isinstance(attributes, list) and any(
|
||||
isinstance(attribute, dict) and attribute.get("name") == "shcoef"
|
||||
for attribute in attributes
|
||||
)
|
||||
if has_sh:
|
||||
members.add(related("shcoef.bin"))
|
||||
environment = related("environment.bin")
|
||||
if (root / Path(*environment.split("/"))).exists():
|
||||
members.add(environment)
|
||||
return members
|
||||
|
||||
root_node = document.get("root")
|
||||
if not isinstance(root_node, dict):
|
||||
raise GaussianPipelineIntegrityError("Gaussian LCC2 descriptor has no root object")
|
||||
splat_files: object
|
||||
if all(key in document for key in ("total_splats", "lod_3dgs_info", "lod_level")):
|
||||
legacy_files = root_node.get("files")
|
||||
if not isinstance(legacy_files, list):
|
||||
raise GaussianPipelineIntegrityError("Gaussian legacy LCC2 descriptor has no files")
|
||||
normalized: list[str] = []
|
||||
for item in legacy_files:
|
||||
if not isinstance(item, str):
|
||||
raise GaussianPipelineIntegrityError("Gaussian LCC2 chunk path is invalid")
|
||||
value = item[1:] if item.startswith("/") else item
|
||||
normalized.append(value if value.endswith(".sog") else f"{value}.sog")
|
||||
splat_files = normalized
|
||||
else:
|
||||
splat_files = root_node.get("splatFiles")
|
||||
if not isinstance(splat_files, list) or not splat_files:
|
||||
raise GaussianPipelineIntegrityError("Gaussian LCC2 descriptor has no splat files")
|
||||
for item in splat_files:
|
||||
if not isinstance(item, str):
|
||||
raise GaussianPipelineIntegrityError("Gaussian LCC2 chunk path is invalid")
|
||||
members.add(related(item))
|
||||
return members
|
||||
|
||||
|
||||
def _logical_path(value: str, label: str) -> str:
|
||||
if (
|
||||
not value
|
||||
or len(value) > 1024
|
||||
or value.startswith("/")
|
||||
or "\\" in value
|
||||
or "\x00" in value
|
||||
or any(part in {"", ".", ".."} for part in value.split("/"))
|
||||
):
|
||||
raise GaussianPipelineIntegrityError(f"Gaussian source {label} is unsafe")
|
||||
return value
|
||||
|
||||
|
||||
def _bundle_member(root: Path, logical_path: str) -> Path:
|
||||
logical = _logical_path(logical_path, "bundle member path")
|
||||
candidate = root
|
||||
for part in logical.split("/"):
|
||||
candidate = candidate / part
|
||||
if candidate.is_symlink():
|
||||
raise GaussianPipelineIntegrityError(
|
||||
f"Gaussian source bundle contains a symlink: {logical}"
|
||||
)
|
||||
try:
|
||||
resolved = candidate.resolve(strict=True)
|
||||
resolved.relative_to(root)
|
||||
except (OSError, ValueError) as exc:
|
||||
raise GaussianPipelineIntegrityError(
|
||||
f"Gaussian source bundle member is unavailable: {logical}"
|
||||
) from exc
|
||||
if not resolved.is_file():
|
||||
raise GaussianPipelineIntegrityError(
|
||||
f"Gaussian source bundle member is not a regular file: {logical}"
|
||||
)
|
||||
return resolved
|
||||
|
||||
|
||||
def _offset(value: str | None, byte_length: int) -> int:
|
||||
if value is None or not value.isdigit():
|
||||
raise GaussianPipelineGatewayError("Gaussian upload offset is invalid")
|
||||
|
||||
Reference in New Issue
Block a user