Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 83 additions & 25 deletions src/rgit/doctor.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import hashlib
import json
from pathlib import Path
from typing import Any
Expand Down Expand Up @@ -145,17 +146,20 @@ def _check_feature_payloads(store, findings: list[dict[str, Any]]) -> None:
fid,
)
continue
try:
payload = _load_json_object(store, digest)
except FileNotFoundError:
_add(
findings,
"error",
"missing_feature_payload_object",
"feature payload_hash does not resolve to an object",
fid,
)
payload_ok, payload_bytes = _check_object_reference(
store,
findings,
digest,
fid,
kind="feature_payload",
reference="feature payload_hash",
load_bytes=True,
)
if not payload_ok:
continue
assert payload_bytes is not None
try:
payload = json.loads(payload_bytes)
except (UnicodeDecodeError, json.JSONDecodeError):
_add(
findings,
Expand Down Expand Up @@ -203,14 +207,14 @@ def _check_run_artifacts(store, findings: list[dict[str, Any]]) -> None:
rid,
)
continue
if not store.objects.path_for(digest).exists():
_add(
findings,
"error",
"missing_run_artifact_object",
"run artifact_hash does not resolve to an object",
rid,
)
_check_object_reference(
store,
findings,
digest,
rid,
kind="run_artifact",
reference="run artifact_hash",
)


def _check_proposals(store, findings: list[dict[str, Any]]) -> None:
Expand All @@ -225,13 +229,14 @@ def _check_proposals(store, findings: list[dict[str, Any]]) -> None:
"proposal has no diff_ref",
pid,
)
elif not store.objects.path_for(diff_ref).exists():
_add(
else:
_check_object_reference(
store,
findings,
"error",
"missing_proposal_diff_object",
"proposal diff_ref does not resolve to an object",
diff_ref,
pid,
kind="proposal_diff",
reference="proposal diff_ref",
)
try:
candidates = json.loads(row["candidates"])
Expand Down Expand Up @@ -343,8 +348,61 @@ def _expected_endpoint(edge_type: str, side: str) -> str:
return "known"


def _load_json_object(store, digest: str) -> Any:
return json.loads(store.objects.get(digest))
def _check_object_reference(
store,
findings: list[dict[str, Any]],
digest: str,
subject: str,
*,
kind: str,
reference: str,
load_bytes: bool = False,
) -> tuple[bool, bytes | None]:
"""Return validity and, when requested, the exact bytes that were hashed."""
if not store.objects.is_valid_digest(digest):
_add(
findings,
"error",
f"invalid_{kind}_reference",
f"{reference} is not a canonical lowercase sha256 digest",
subject,
)
return False, None
try:
if load_bytes:
data = store.objects.get(digest)
matches = hashlib.sha256(data).hexdigest() == digest
else:
data = None
matches = store.objects.verify(digest)
except FileNotFoundError:
_add(
findings,
"error",
f"missing_{kind}_object",
f"{reference} does not resolve to an object",
subject,
)
return False, None
except OSError:
_add(
findings,
"error",
f"unreadable_{kind}_object",
f"{kind.replace('_', ' ')} object cannot be read",
subject,
)
return False, None
if not matches:
_add(
findings,
"error",
f"corrupt_{kind}_object",
f"{kind.replace('_', ' ')} object does not match sha256 {digest}",
subject,
)
return False, None
return True, data


def _add(
Expand Down
83 changes: 79 additions & 4 deletions src/rgit/store/objects.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,22 @@
import hashlib
import json
import os
import time
from pathlib import Path
from typing import Any
from typing import Any, Callable, TypeVar
from uuid import uuid4


_T = TypeVar("_T")


class ObjectStore:
"""Immutable sha256-addressed blob store under a directory."""

_DIGEST_LENGTH = hashlib.sha256().digest_size * 2
_DIGEST_CHARS = frozenset("0123456789abcdef")
_ACCESS_RETRY_DELAYS = (0.001, 0.002, 0.004, 0.008, 0.016, 0.032)

def __init__(self, root: Path, create: bool = True):
self.root = Path(root)
if create:
Expand All @@ -15,17 +25,82 @@ def __init__(self, root: Path, create: bool = True):
def path_for(self, digest: str) -> Path:
"""On-disk location of a digest — the single source of truth for the
store layout, so read-only consumers (doctor) can't drift from it."""
if not self.is_valid_digest(digest):
raise ValueError(f"invalid sha256 digest: {digest!r}")
return self.root / digest[:2] / digest[2:]

def _path(self, digest: str) -> Path:
return self.path_for(digest)

@classmethod
def is_valid_digest(cls, digest: object) -> bool:
"""Return whether *digest* is a canonical lowercase sha256 hex value."""
return (
isinstance(digest, str)
and len(digest) == cls._DIGEST_LENGTH
and all(char in cls._DIGEST_CHARS for char in digest)
)

def verify(self, digest: str) -> bool:
"""Return whether the stored bytes hash to *digest*.

Missing objects still raise ``FileNotFoundError`` so callers such as
doctor can distinguish an absent object from a corrupt one.
"""
path = self.path_for(digest)
with path.open("rb") as blob:
actual = hashlib.file_digest(blob, "sha256").hexdigest()
return actual == digest

@classmethod
def _retry_permission_error(cls, operation: Callable[[], _T]) -> _T:
"""Retry brief sharing violations without hiding persistent errors."""
for delay in cls._ACCESS_RETRY_DELAYS:
try:
return operation()
except PermissionError:
time.sleep(delay)
return operation()

@classmethod
def _write_atomic(cls, path: Path, data: bytes) -> None:
"""Durably write *data* before atomically publishing it at *path*."""
# Path.open("x") preserves the mode/umask behavior of the previous
# direct write while guaranteeing that a stale temp file is not reused.
temp_path = path.with_name(f".{path.name}.{uuid4().hex}.tmp")
try:
with temp_path.open("xb") as blob:
blob.write(data)
blob.flush()
os.fsync(blob.fileno())
cls._retry_permission_error(lambda: os.replace(temp_path, path))
finally:
try:
temp_path.unlink()
except OSError:
pass

def put(self, data: bytes) -> str:
digest = hashlib.sha256(data).hexdigest()
p = self._path(digest)
if not p.exists():
p.parent.mkdir(parents=True, exist_ok=True)
p.write_bytes(data)
try:
if self._retry_permission_error(lambda: self.verify(digest)):
return digest
except FileNotFoundError:
pass

p.parent.mkdir(parents=True, exist_ok=True)
try:
self._write_atomic(p, data)
except OSError:
# Another writer may have won the race. This is success only when
# it published the exact object we were trying to store.
try:
if self._retry_permission_error(lambda: self.verify(digest)):
return digest
except OSError:
pass
raise
return digest

def get(self, digest: str) -> bytes:
Expand Down
117 changes: 117 additions & 0 deletions tests/test_doctor.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,65 @@ def test_doctor_reports_missing_feature_payload_object(git_repo):
assert "missing_feature_payload_object" in _codes(report, level="error")


def test_doctor_reports_corrupt_feature_payload_object(git_repo):
from rgit.doctor import run_doctor

store = Store.init(git_repo)
fid = store.add_feature(_cap())
payload_hash = store.conn.execute(
"SELECT payload_hash FROM features WHERE id=?", (fid,)
).fetchone()["payload_hash"]
_object_path(store, payload_hash).write_bytes(b"[]")

report = run_doctor(store)

assert report["ok"] is False
assert "corrupt_feature_payload_object" in _codes(report, level="error")
assert "malformed_feature_payload_json" not in _codes(report)


def test_doctor_hashes_the_same_feature_payload_bytes_it_parses(
git_repo, monkeypatch
):
from rgit.doctor import run_doctor

store = Store.init(git_repo)
fid = store.add_feature(_cap())
payload_hash = store.conn.execute(
"SELECT payload_hash FROM features WHERE id=?", (fid,)
).fetchone()["payload_hash"]
real_get = store.objects.get

def tamper_before_read(digest):
_object_path(store, payload_hash).write_bytes(b"[]")
return real_get(digest)

monkeypatch.setattr(store.objects, "get", tamper_before_read)

report = run_doctor(store)

assert report["ok"] is False
assert "corrupt_feature_payload_object" in _codes(report, level="error")
assert "malformed_feature_payload_json" not in _codes(report)


def test_doctor_reports_feature_payload_read_race(git_repo, monkeypatch):
from rgit.doctor import run_doctor

store = Store.init(git_repo)
store.add_feature(_cap())

def deny_read(digest):
raise PermissionError("simulated read race")

monkeypatch.setattr(store.objects, "get", deny_read)

report = run_doctor(store)

assert report["ok"] is False
assert "unreadable_feature_payload_object" in _codes(report, level="error")


def test_doctor_reports_missing_run_artifact_object(git_repo):
from rgit.doctor import run_doctor

Expand All @@ -116,6 +175,20 @@ def test_doctor_reports_missing_run_artifact_object(git_repo):
assert "missing_run_artifact_object" in _codes(report, level="error")


def test_doctor_reports_corrupt_run_artifact_object(git_repo):
from rgit.doctor import run_doctor

store = Store.init(git_repo)
artifact_hash = store.objects.put(b"artifact")
_run(store, artifact=artifact_hash)
_object_path(store, artifact_hash).write_bytes(b"tampered")

report = run_doctor(store)

assert report["ok"] is False
assert "corrupt_run_artifact_object" in _codes(report, level="error")


def test_doctor_reports_missing_proposal_diff_object(git_repo):
from rgit.doctor import run_doctor

Expand All @@ -129,6 +202,50 @@ def test_doctor_reports_missing_proposal_diff_object(git_repo):
assert "missing_proposal_diff_object" in _codes(report, level="error")


def test_doctor_reports_corrupt_proposal_diff_object(git_repo):
from rgit.doctor import run_doctor

store = Store.init(git_repo)
diff_ref = store.objects.put(b"diff")
_proposal(store, diff=diff_ref)
_object_path(store, diff_ref).write_bytes(b"tampered")

report = run_doctor(store)

assert report["ok"] is False
assert "corrupt_proposal_diff_object" in _codes(report, level="error")


def test_doctor_reports_invalid_object_reference_without_path_lookup(git_repo):
from rgit.doctor import run_doctor

store = Store.init(git_repo)
_run(store, artifact="../outside")

report = run_doctor(store)

assert report["ok"] is False
assert "invalid_run_artifact_reference" in _codes(report, level="error")


def test_doctor_reports_unreadable_object(git_repo, monkeypatch):
from rgit.doctor import run_doctor

store = Store.init(git_repo)
artifact_hash = store.objects.put(b"artifact")
_run(store, artifact=artifact_hash)

def deny_read(digest):
raise PermissionError("simulated unreadable object")

monkeypatch.setattr(store.objects, "verify", deny_read)

report = run_doctor(store)

assert report["ok"] is False
assert "unreadable_run_artifact_object" in _codes(report, level="error")


def test_doctor_reports_malformed_proposal_candidates_json(git_repo):
from rgit.doctor import run_doctor

Expand Down
Loading
Loading