fix: complete durable workspace runtime config handoff
This commit is contained in:
+74
-12
@@ -1,3 +1,4 @@
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
@@ -572,25 +573,80 @@ def _validate_raw_config_shape(raw: dict[str, Any], path: Path) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _read_runtime_fd(fd: int, label: str, expected_mode: int = 0o400) -> tuple[bytes, os.stat_result]:
|
||||
try:
|
||||
info = os.fstat(fd)
|
||||
if (not stat.S_ISREG(info.st_mode) or info.st_nlink != 1
|
||||
or stat.S_IMODE(info.st_mode) != expected_mode
|
||||
or info.st_uid != os.getuid()):
|
||||
raise OSError("unsafe runtime descriptor")
|
||||
os.lseek(fd, 0, os.SEEK_SET)
|
||||
chunks: list[bytes] = []
|
||||
total = 0
|
||||
while chunk := os.read(fd, 1024 * 1024):
|
||||
total += len(chunk)
|
||||
if total > 16 * 1024 * 1024:
|
||||
raise OSError("runtime descriptor too large")
|
||||
chunks.append(chunk)
|
||||
return b"".join(chunks), info
|
||||
except OSError as exc:
|
||||
raise ConfigError(f"File runtime {label} non attendibile") from exc
|
||||
|
||||
|
||||
def _strict_runtime_manifest(raw: object) -> dict[str, object]:
|
||||
required = {
|
||||
"version", "workspace_id", "workspace_revision", "descriptor_git_blob",
|
||||
"descriptor_sha256", "config_sha256", "config_dwh_binding", "config_dev",
|
||||
"config_ino", "config_size", "config_mode", "config_uid", "config_nlink",
|
||||
}
|
||||
if not isinstance(raw, dict) or set(raw) != required or raw.get("version") != 1:
|
||||
raise ConfigError("Manifest runtime non valido")
|
||||
if not isinstance(raw.get("config_dwh_binding"), dict):
|
||||
raise ConfigError("Manifest runtime non valido")
|
||||
binding = raw["config_dwh_binding"]
|
||||
if set(binding) != {"workspace_id", "config_fingerprint", "input_fingerprint"} or any(not isinstance(v, str) for v in binding.values()):
|
||||
raise ConfigError("Manifest runtime non valido")
|
||||
for key in ("descriptor_sha256", "config_sha256"):
|
||||
if not isinstance(raw[key], str) or not re.fullmatch(r"[0-9a-f]{64}", raw[key]):
|
||||
raise ConfigError("Manifest runtime non valido")
|
||||
for key in ("config_dev", "config_ino", "config_size", "config_uid", "config_nlink"):
|
||||
if not isinstance(raw[key], str) or not raw[key].isdigit():
|
||||
raise ConfigError("Manifest runtime non valido")
|
||||
if raw["config_mode"] != "400":
|
||||
raise ConfigError("Manifest runtime non valido")
|
||||
return raw
|
||||
|
||||
|
||||
def load_config(path: Path) -> Config:
|
||||
# Backend runtime leases pass the verified canonical config as fd 3 while retaining
|
||||
# the ordinary absolute -c argument for diagnostics and source identity. Never reopen
|
||||
# that pathname: an ancestor or leaf replacement after spawn must not alter bytes used
|
||||
# by the harness.
|
||||
runtime_manifest: dict[str, object] | None = None
|
||||
runtime_fd = os.environ.get("THT_CONFIG_FD")
|
||||
if runtime_fd is not None:
|
||||
manifest_fd = os.environ.get("THT_CONFIG_MANIFEST_FD")
|
||||
expected_manifest = os.environ.get("THT_CONFIG_MANIFEST_SHA256")
|
||||
if runtime_fd is not None or manifest_fd is not None or expected_manifest is not None:
|
||||
if runtime_fd is None or manifest_fd is None or expected_manifest is None or not re.fullmatch(r"[0-9a-f]{64}", expected_manifest):
|
||||
raise ConfigError("Handoff runtime incompleto")
|
||||
try:
|
||||
fd = int(runtime_fd)
|
||||
info = os.fstat(fd)
|
||||
if (not stat.S_ISREG(info.st_mode) or info.st_nlink != 1
|
||||
or stat.S_IMODE(info.st_mode) != 0o400
|
||||
or info.st_uid != os.getuid()):
|
||||
raise OSError("unsafe runtime config descriptor")
|
||||
chunks: list[bytes] = []
|
||||
while chunk := os.read(fd, 1024 * 1024):
|
||||
chunks.append(chunk)
|
||||
source_text = b"".join(chunks).decode("utf-8")
|
||||
except (OSError, UnicodeError, ValueError) as exc:
|
||||
config_bytes, config_info = _read_runtime_fd(int(runtime_fd), "config")
|
||||
manifest_bytes, manifest_info = _read_runtime_fd(int(manifest_fd), "manifest", 0o600)
|
||||
if hashlib.sha256(manifest_bytes).hexdigest() != expected_manifest:
|
||||
raise ConfigError("Manifest runtime modificato")
|
||||
runtime_manifest = _strict_runtime_manifest(json.loads(manifest_bytes.decode("utf-8")))
|
||||
if (runtime_manifest["config_sha256"] != hashlib.sha256(config_bytes).hexdigest()
|
||||
or int(runtime_manifest["config_dev"]) != config_info.st_dev
|
||||
or int(runtime_manifest["config_ino"]) != config_info.st_ino
|
||||
or int(runtime_manifest["config_size"]) != config_info.st_size
|
||||
or int(runtime_manifest["config_uid"]) != config_info.st_uid
|
||||
or int(runtime_manifest["config_nlink"]) != config_info.st_nlink
|
||||
or config_info.st_dev == manifest_info.st_dev and config_info.st_ino == manifest_info.st_ino):
|
||||
raise ConfigError("Identità config runtime non valida")
|
||||
source_text = config_bytes.decode("utf-8")
|
||||
except (OSError, UnicodeError, ValueError, json.JSONDecodeError) as exc:
|
||||
if isinstance(exc, ConfigError):
|
||||
raise
|
||||
raise ConfigError("File di configurazione runtime non attendibile") from exc
|
||||
else:
|
||||
if not path.exists():
|
||||
@@ -679,6 +735,12 @@ def load_config(path: Path) -> Config:
|
||||
)
|
||||
_validate_active_embeddings_config(cfg.embeddings, path)
|
||||
_validate_active_vector_config(cfg.vectors, path)
|
||||
if runtime_manifest is not None:
|
||||
from tht.jobs.dwh_pipeline import config_dwh_binding
|
||||
if runtime_manifest["workspace_id"] != cfg._workspace_id or runtime_manifest["workspace_revision"] != cfg._workspace_revision:
|
||||
raise ConfigError("Identità workspace runtime non valida")
|
||||
if config_dwh_binding(cfg) != runtime_manifest["config_dwh_binding"]:
|
||||
raise ConfigError("Binding DWH runtime modificato")
|
||||
return cfg
|
||||
|
||||
|
||||
|
||||
@@ -26,8 +26,16 @@ def safe_rev(v: str) -> bool:
|
||||
|
||||
|
||||
def open_dir(parent: int | None, name: str, create: bool = False) -> int:
|
||||
flags = os.O_RDONLY | getattr(os, "O_DIRECTORY", 0) | os.O_NOFOLLOW
|
||||
# Darwin rejects O_NOFOLLOW|openat for directories (ELOOP); lstat the
|
||||
# component before opening and verify the resulting descriptor below. Linux
|
||||
# uses the stronger flag where available.
|
||||
flags = os.O_RDONLY | getattr(os, "O_DIRECTORY", 0)
|
||||
if sys.platform != "darwin":
|
||||
flags |= os.O_NOFOLLOW
|
||||
try:
|
||||
entry = os.stat(name, dir_fd=parent, follow_symlinks=False)
|
||||
if stat.S_ISLNK(entry.st_mode):
|
||||
fail("runtime config directory is not trusted")
|
||||
return os.open(name, flags, dir_fd=parent)
|
||||
except FileNotFoundError:
|
||||
if not create:
|
||||
@@ -48,25 +56,39 @@ def checked_dir(fd: int, expected_mode: int = 0o700) -> None:
|
||||
|
||||
|
||||
def walk(root: str, comps: list[str], create: bool = True) -> int:
|
||||
"""Open an absolute path component-by-component without following symlinks.
|
||||
|
||||
In particular, never use os.makedirs/root pathname resolution here: an attacker
|
||||
replacing an ancestor between those calls must not redirect publication.
|
||||
"""
|
||||
if not os.path.isabs(root):
|
||||
fail("data root must be absolute")
|
||||
if not os.path.lexists(root):
|
||||
os.makedirs(root, mode=0o700, exist_ok=True)
|
||||
fd = os.open(root, os.O_RDONLY | getattr(os, "O_DIRECTORY", 0) | os.O_NOFOLLOW)
|
||||
# macOS exposes temporary directories through the conventional /var and
|
||||
# /tmp symlinks. Resolve only these OS-owned aliases; workspace-owned
|
||||
# ancestors remain component checked and are never realpath-followed.
|
||||
if root == "/var" or root == "/tmp" or root.startswith(("/var/", "/tmp/")):
|
||||
root = "/private" + root
|
||||
parts = [part for part in Path(root).parts if part not in ("", "/")]
|
||||
if any(part in (".", "..") or "/" in part for part in parts + comps):
|
||||
fail("unsafe path component")
|
||||
fd = os.open("/", os.O_RDONLY | getattr(os, "O_DIRECTORY", 0))
|
||||
try:
|
||||
checked_dir(fd)
|
||||
except:
|
||||
all_components = [*parts, *comps]
|
||||
for index, component in enumerate(all_components):
|
||||
nxt = open_dir(fd, component, create)
|
||||
# Ancestors such as /var/folders are installation-owned and commonly
|
||||
# 0755; the trusted runtime root and every workspace child are private.
|
||||
info = os.fstat(nxt)
|
||||
if (not stat.S_ISDIR(info.st_mode) or info.st_nlink < 1
|
||||
or (index >= len(parts) - 1 and (info.st_uid != os.getuid() or stat.S_IMODE(info.st_mode) != 0o700))):
|
||||
os.close(nxt)
|
||||
fail("runtime config directory is not trusted")
|
||||
os.close(fd)
|
||||
fd = nxt
|
||||
return fd
|
||||
except BaseException:
|
||||
os.close(fd)
|
||||
raise
|
||||
for c in comps:
|
||||
if c in ("", ".", "..") or "/" in c:
|
||||
os.close(fd)
|
||||
fail("unsafe path component")
|
||||
nxt = open_dir(fd, c, create)
|
||||
checked_dir(nxt)
|
||||
os.close(fd)
|
||||
fd = nxt
|
||||
return fd
|
||||
|
||||
|
||||
def read_regular(fd: int, mode: int, expected: bytes | None = None) -> os.stat_result:
|
||||
@@ -100,6 +122,48 @@ def write_all(fd: int, data: bytes) -> None:
|
||||
pos += n
|
||||
|
||||
|
||||
def read_all(fd: int, limit: int = 16 * 1024 * 1024) -> bytes:
|
||||
os.lseek(fd, 0, os.SEEK_SET)
|
||||
chunks: list[bytes] = []
|
||||
total = 0
|
||||
while True:
|
||||
chunk = os.read(fd, min(1024 * 1024, limit - total))
|
||||
if not chunk:
|
||||
return b"".join(chunks)
|
||||
chunks.append(chunk)
|
||||
total += len(chunk)
|
||||
if total > limit:
|
||||
fail("runtime config file is too large")
|
||||
|
||||
|
||||
def strict_manifest(value: object) -> dict:
|
||||
if not isinstance(value, dict):
|
||||
fail("runtime config manifest is invalid")
|
||||
required = {
|
||||
"version", "workspace_id", "workspace_revision", "descriptor_git_blob",
|
||||
"descriptor_sha256", "config_sha256", "config_dwh_binding", "config_dev",
|
||||
"config_ino", "config_size", "config_mode", "config_uid", "config_nlink",
|
||||
}
|
||||
if set(value) != required or value.get("version") != 1:
|
||||
fail("runtime config manifest is invalid")
|
||||
if not safe_id(value.get("workspace_id")) or not safe_rev(value.get("workspace_revision")):
|
||||
fail("runtime config manifest is invalid")
|
||||
if not isinstance(value.get("descriptor_git_blob"), str) or not safe_rev(value["descriptor_git_blob"]):
|
||||
fail("runtime config manifest is invalid")
|
||||
for key in ("descriptor_sha256", "config_sha256"):
|
||||
if not isinstance(value[key], str) or not __import__("re").fullmatch(r"[0-9a-f]{64}", value[key]):
|
||||
fail("runtime config manifest is invalid")
|
||||
binding_value = value.get("config_dwh_binding")
|
||||
if not isinstance(binding_value, dict) or set(binding_value) != {"workspace_id", "config_fingerprint", "input_fingerprint"} or any(not isinstance(x, str) for x in binding_value.values()):
|
||||
fail("runtime config manifest is invalid")
|
||||
for key in ("config_dev", "config_ino", "config_size", "config_uid", "config_nlink"):
|
||||
if not isinstance(value[key], str) or not value[key].isdigit():
|
||||
fail("runtime config manifest is invalid")
|
||||
if value["config_mode"] != "400":
|
||||
fail("runtime config manifest is invalid")
|
||||
return value
|
||||
|
||||
|
||||
def publish(inp: dict) -> dict:
|
||||
root = inp.get("data_root")
|
||||
wid = inp.get("workspace_id")
|
||||
@@ -160,7 +224,7 @@ def publish(inp: dict) -> dict:
|
||||
if got:
|
||||
fd, s = got
|
||||
os.lseek(fd, 0, os.SEEK_SET)
|
||||
old = os.read(fd, len(content) + 1)
|
||||
old = read_all(fd)
|
||||
os.close(fd)
|
||||
if old != content:
|
||||
fail("same-revision runtime configuration changed")
|
||||
@@ -179,6 +243,10 @@ def publish(inp: dict) -> dict:
|
||||
)
|
||||
except FileExistsError:
|
||||
pass
|
||||
# Keep metadata ordering explicit even on filesystems where a
|
||||
# hardlink publication does not retain fchmod as expected.
|
||||
os.fchmod(fd, 0o400)
|
||||
os.fsync(fd)
|
||||
finally:
|
||||
os.close(fd)
|
||||
try:
|
||||
@@ -191,7 +259,7 @@ def publish(inp: dict) -> dict:
|
||||
fd, s = got
|
||||
try:
|
||||
os.lseek(fd, 0, os.SEEK_SET)
|
||||
if os.read(fd, len(content) + 1) != content:
|
||||
if read_all(fd) != content:
|
||||
fail("same-revision runtime configuration changed")
|
||||
finally:
|
||||
os.close(fd)
|
||||
@@ -201,7 +269,7 @@ def publish(inp: dict) -> dict:
|
||||
assert got
|
||||
fd, s = got
|
||||
os.close(fd)
|
||||
manifest = dict(base)
|
||||
manifest = {"version": 1, **dict(base)}
|
||||
manifest.update(
|
||||
{
|
||||
"config_sha256": hashlib.sha256(content).hexdigest(),
|
||||
@@ -213,13 +281,16 @@ def publish(inp: dict) -> dict:
|
||||
"config_nlink": str(s.st_nlink),
|
||||
}
|
||||
)
|
||||
strict_manifest(manifest)
|
||||
mb = (json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n").encode()
|
||||
oldm = current(mandir, mname, 0o600)
|
||||
if oldm:
|
||||
mfd, _ = oldm
|
||||
os.lseek(mfd, 0, os.SEEK_SET)
|
||||
existing = os.read(mfd, len(mb) + 1)
|
||||
existing = read_all(mfd)
|
||||
os.close(mfd)
|
||||
try: strict_manifest(json.loads(existing.decode()))
|
||||
except (ValueError, TypeError, UnicodeError, RuntimeError): fail("runtime config manifest is invalid")
|
||||
if existing != mb:
|
||||
fail("same-revision runtime configuration changed")
|
||||
else:
|
||||
@@ -245,6 +316,7 @@ def publish(inp: dict) -> dict:
|
||||
"path": f"{root}/sessions/{wid}/preprocessing/runtime-config/{name}",
|
||||
"manifestPath": f"{root}/sessions/{wid}/preprocessing/runtime-config-manifests/{mname}",
|
||||
"manifest": mb.decode(),
|
||||
"manifest_sha256": hashlib.sha256(mb).hexdigest(),
|
||||
"dev": s.st_dev,
|
||||
"ino": s.st_ino,
|
||||
}
|
||||
@@ -295,35 +367,53 @@ def verified_snapshot(inp: dict) -> dict:
|
||||
payload += x
|
||||
finally:
|
||||
os.close(mf)
|
||||
manifest = json.loads(payload.decode())
|
||||
record = next((r for r in manifest.get("revisions", []) if r.get("id") == wid), None)
|
||||
try:
|
||||
manifest = json.loads(payload.decode())
|
||||
except (UnicodeDecodeError, json.JSONDecodeError):
|
||||
fail("workspace snapshot integrity check failed")
|
||||
if not isinstance(manifest, dict) or set(manifest) != {"head", "revisions", "files"}:
|
||||
fail("workspace snapshot integrity check failed")
|
||||
records = manifest.get("revisions")
|
||||
files = manifest.get("files")
|
||||
record = next((r for r in records if isinstance(r, dict) and r.get("id") == wid), None) if isinstance(records, list) else None
|
||||
expected_path = f"{root}/{rev}/{wid}.yaml"
|
||||
if (
|
||||
manifest.get("head") != rev
|
||||
or not record
|
||||
or record.get("commit") != rev
|
||||
or manifest.get("files", {}).get(f"{wid}.yaml") != hashlib.sha256(source).hexdigest()
|
||||
manifest.get("head") != rev or not isinstance(records, list) or not record
|
||||
or set(record) != {"id", "commit", "blob", "snapshotPath"}
|
||||
or record.get("commit") != rev or record.get("snapshotPath") != expected_path
|
||||
or not isinstance(record.get("blob"), str) or not safe_rev(record.get("blob"))
|
||||
or not isinstance(files, dict)
|
||||
or files.get(f"{wid}.yaml") != hashlib.sha256(source).hexdigest()
|
||||
):
|
||||
fail("workspace snapshot integrity check failed")
|
||||
repo = inp.get("repository_root")
|
||||
if repo:
|
||||
if not isinstance(repo, str) or not os.path.isabs(repo):
|
||||
fail("invalid repository root")
|
||||
try:
|
||||
blob = subprocess.check_output(
|
||||
["git", "-C", repo, "rev-parse", f"{rev}:workspaces/{wid}.yaml"],
|
||||
stderr=subprocess.DEVNULL,
|
||||
text=True,
|
||||
timeout=5,
|
||||
).strip()
|
||||
git_source = subprocess.check_output(
|
||||
["git", "-C", repo, "show", f"{rev}:workspaces/{wid}.yaml"],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=5,
|
||||
)
|
||||
except (OSError, subprocess.SubprocessError):
|
||||
fail("workspace Git revision is unavailable")
|
||||
if blob != record.get("blob") or git_source != source:
|
||||
fail("workspace Git descriptor identity mismatch")
|
||||
if not isinstance(repo, str) or not os.path.isabs(repo):
|
||||
fail("invalid repository root")
|
||||
try:
|
||||
blob = subprocess.check_output(
|
||||
["git", "-C", repo, "rev-parse", f"{rev}:workspaces/{wid}.yaml"],
|
||||
stderr=subprocess.DEVNULL, text=True, timeout=5,
|
||||
).strip()
|
||||
git_source = subprocess.check_output(
|
||||
["git", "-C", repo, "show", f"{rev}:workspaces/{wid}.yaml"],
|
||||
stderr=subprocess.DEVNULL, timeout=5,
|
||||
)
|
||||
except (OSError, subprocess.SubprocessError):
|
||||
fail("workspace Git revision is unavailable")
|
||||
try:
|
||||
import re
|
||||
from collections import Counter
|
||||
normalize = lambda value: re.findall(r"[A-Za-z0-9_.:/@+-]+", value)
|
||||
git_tokens = Counter(normalize(git_source.decode("utf-8")))
|
||||
snapshot_tokens = Counter(normalize(source.decode("utf-8")))
|
||||
# The registry canonicalizer may add schema defaults/reorder mappings.
|
||||
# Every token from the exact Git descriptor must nevertheless survive;
|
||||
# replacements (including non-rendered workspace.name) are rejected.
|
||||
equivalent = all(snapshot_tokens[k] >= count for k, count in git_tokens.items())
|
||||
except UnicodeDecodeError:
|
||||
equivalent = False
|
||||
if blob != record.get("blob") or not equivalent:
|
||||
fail("workspace Git descriptor identity mismatch")
|
||||
return {
|
||||
"workspace_id": wid,
|
||||
"workspace_revision": rev,
|
||||
|
||||
Reference in New Issue
Block a user