Skip to content
Open
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
22 changes: 14 additions & 8 deletions minigit/index.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
"""

import os
import stat
from dataclasses import dataclass
from typing import NamedTuple

Expand Down Expand Up @@ -41,6 +42,14 @@ def __init__(self, repo_path=".", store=None):

self.store = store

@staticmethod
def _file_mode(path: str) -> str:
if os.stat(path).st_mode & stat.S_IXUSR:
return "100755"
if os.name == "nt" and os.path.splitext(path)[1].lower() in {".bat", ".cmd", ".exe", ".sh"}:
return "100755"
return "100644"

def read_index(self) -> list[IndexEntry]:
if not os.path.exists(self.index_path):
return []
Expand Down Expand Up @@ -81,10 +90,7 @@ def stage_file(self, path) -> None:

blob_hash = self.store.write_object(data, "blob")

if os.access(full_path, os.X_OK):
mode = "100755"
else:
mode = "100644"
mode = self._file_mode(full_path)

entries = [e for e in self.read_index() if e.path != rel_path]
entries.append(IndexEntry(mode, blob_hash, rel_path))
Expand Down Expand Up @@ -147,7 +153,7 @@ def diff_working_tree_vs(self, tree_hash) -> DiffResult:
with open(full_path, "rb") as f:
data = f.read()
current_hash = self.store.hash_object(data, "blob")
current_mode = "100755" if os.access(full_path, os.X_OK) else "100644"
current_mode = self._file_mode(full_path)
if current_hash != entry.hash or current_mode != entry.mode:
result.modified.append(entry.path)

Expand Down Expand Up @@ -177,7 +183,7 @@ def _working_status(self) -> DiffResult:
with open(full_path, "rb") as f:
data = f.read()
current_hash = self.store.hash_object(data, "blob")
current_mode = "100755" if os.access(full_path, os.X_OK) else "100644"
current_mode = self._file_mode(full_path)
if current_hash != entry.hash or current_mode != entry.mode:
result.modified.append(entry.path)

Expand Down Expand Up @@ -216,7 +222,7 @@ def checkout(self, tree_hash) -> None:
with open(full_path, "rb") as f:
data = f.read()
disk_hash = self.store.hash_object(data, "blob")
disk_mode = "100755" if os.access(full_path, os.X_OK) else "100644"
disk_mode = self._file_mode(full_path)
if disk_hash != entry.hash or disk_mode != entry.mode:
raise MiniGitError(f"local changes would be lost: {entry.path}")

Expand Down Expand Up @@ -335,7 +341,7 @@ def cmd_status(args) -> int:
with open(full_path, "rb") as f:
data = f.read()
current_hash = wt.store.hash_object(data, "blob")
current_mode = "100755" if os.access(full_path, os.X_OK) else "100644"
current_mode = wt._file_mode(full_path)
if current_hash != entry.hash or current_mode != entry.mode:
not_staged_paths.add(entry.path)

Expand Down
226 changes: 216 additions & 10 deletions minigit/objects.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,11 @@

import hashlib
import os
import struct
import tempfile
import uuid
import zlib
from collections.abc import Iterable
from pathlib import Path
from typing import NamedTuple

Expand All @@ -28,6 +31,13 @@ class TreeEntry(NamedTuple):

class ObjectStore:
_OBJECT_TYPES = {"blob", "tree", "commit"}
_PACK_MAGIC = b"MGPK"
_INDEX_MAGIC = b"MGIX"
_PACK_VERSION = 1
_PACK_HEADER = struct.Struct(">4sII")
_RECORD_HEADER = struct.Struct(">BQQ")
_INDEX_HEADER = struct.Struct(">4sII")
_INDEX_ENTRY = struct.Struct(">20sQQ")

root: Path
objects_dir: Path
Expand All @@ -38,6 +48,7 @@ def __init__(self, repo_path=".") -> None:
"""
self.root = Path(repo_path)
self.objects_dir = self.root / ".minigit" / "objects"
self.pack_dir = self.objects_dir / "pack"

def hash_object(self, data: bytes, obj_type: str) -> str:
"""
Expand All @@ -61,6 +72,12 @@ def write_object(self, data: bytes, obj_type: str) -> str:
if object_path.exists():
self.read_object(obj_hash)
return obj_hash
try:
self.read_object(obj_hash)
except ObjectNotFoundError:
pass
else:
return obj_hash

header = f"{obj_type} {len(data)}".encode()
object_path.parent.mkdir(parents=True, exist_ok=True)
Expand All @@ -87,41 +104,230 @@ def read_object(self, hash: str) -> tuple[str, bytes]:
ObjectNotFoundError for invalid or missing hashes and
ObjectCorruptError for invalid compressed data, headers, lengths,
types, or content hashes.
Perform a loose-first lookup, then a packed lookup, validating every object
"""
if len(hash) != 40 or any(character not in "0123456789abcdef" for character in hash):
raise ObjectNotFoundError(hash)

object_path = self._object_path(hash)

if not object_path.exists():
if object_path.exists():
return self._read_loose_object(hash, object_path)

packed = self._read_packed_object(hash)
if packed is None:
raise ObjectNotFoundError(hash)
return packed

def pack_objects(self, hashes: Iterable[str]) -> tuple[Path, Path]:
"""Pack selected objects and atomically publish a pack/index pair.
Reads through the public API to validate and deduplicate objects, then writes
a pack file and an index file. Returns the paths of the pack and index files.
"""
selected: dict[str, tuple[str, bytes]] = {}
for obj_hash in hashes:
if obj_hash not in selected:
selected[obj_hash] = self.read_object(obj_hash)

self.pack_dir.mkdir(parents=True, exist_ok=True)
pack_id = uuid.uuid4().hex
pack_path = self.pack_dir / f"pack-{pack_id}.pack"
index_path = self.pack_dir / f"pack-{pack_id}.idx"
pack_temp: Path | None = None
index_temp: Path | None = None
entries: list[tuple[str, int, int]] = []

try:
with tempfile.NamedTemporaryFile(
dir=self.pack_dir, prefix=f".pack-{pack_id}.", suffix=".tmp", delete=False
) as pack_file:
pack_temp = Path(pack_file.name)
pack_file.write(
self._PACK_HEADER.pack(self._PACK_MAGIC, self._PACK_VERSION, len(selected))
)
for obj_hash, (obj_type, data) in selected.items():
offset = pack_file.tell()
compressed = zlib.compress(data)
type_bytes = obj_type.encode("ascii")
record = self._RECORD_HEADER.pack(len(type_bytes), len(data), len(compressed))
pack_file.write(record)
pack_file.write(type_bytes)
pack_file.write(compressed)
record_length = self._RECORD_HEADER.size + len(type_bytes) + len(compressed)
entries.append((obj_hash, offset, record_length))
pack_file.flush()
os.fsync(pack_file.fileno())

self._validate_pack(pack_temp)
os.replace(pack_temp, pack_path)
pack_temp = None

with tempfile.NamedTemporaryFile(
dir=self.pack_dir, prefix=f".pack-{pack_id}.", suffix=".tmp", delete=False
) as index_file:
index_temp = Path(index_file.name)
ordered_entries = sorted(entries, key=lambda entry: entry[0])
index_file.write(
self._INDEX_HEADER.pack(self._INDEX_MAGIC, self._PACK_VERSION, len(entries))
)
for obj_hash, offset, length in ordered_entries:
index_file.write(
self._INDEX_ENTRY.pack(bytes.fromhex(obj_hash), offset, length)
)
index_file.flush()
os.fsync(index_file.fileno())

self._read_index(index_temp, pack_path)
os.replace(index_temp, index_path)
index_temp = None
finally:
if pack_temp is not None:
pack_temp.unlink(missing_ok=True)
if index_temp is not None:
index_temp.unlink(missing_ok=True)

return pack_path, index_path

def _read_loose_object(self, obj_hash: str, object_path: Path) -> tuple[str, bytes]:
try:
compressed_object = object_path.read_bytes()
decompressor = zlib.decompressobj()
raw_object = decompressor.decompress(compressed_object) + decompressor.flush()
if not decompressor.eof or decompressor.unused_data or decompressor.unconsumed_tail:
raise ObjectCorruptError(hash)
raise ObjectCorruptError(obj_hash)
header, separator, content = raw_object.partition(b"\0")
if not separator:
raise ObjectCorruptError(hash)
raise ObjectCorruptError(obj_hash)

header_fields = header.split(b" ")
if len(header_fields) != 2:
raise ObjectCorruptError(hash)
raise ObjectCorruptError(obj_hash)
obj_type_bytes, obj_length_bytes = header_fields
obj_type = obj_type_bytes.decode("ascii")
if obj_type not in self._OBJECT_TYPES or not obj_length_bytes.isdigit():
raise ObjectCorruptError(hash)
raise ObjectCorruptError(obj_hash)
if int(obj_length_bytes) != len(content):
raise ObjectCorruptError(hash)

raise ObjectCorruptError(obj_hash)
except (OSError, UnicodeDecodeError, ValueError, zlib.error) as error:
raise ObjectCorruptError(hash) from error
raise ObjectCorruptError(obj_hash) from error

if self.hash_object(content, obj_type) != obj_hash:
raise ObjectCorruptError(obj_hash)
return obj_type, content

if self.hash_object(content, obj_type) != hash:
raise ObjectCorruptError(hash)
def _read_packed_object(self, obj_hash: str) -> tuple[str, bytes] | None:
if not self.pack_dir.exists():
return None
for index_path in sorted(self.pack_dir.glob("pack-*.idx")):
entries, pack_path = self._read_index(index_path)
entry = entries.get(obj_hash)
if entry is None:
continue
offset, record_length = entry
try:
with pack_path.open("rb") as pack_file:
pack_file.seek(offset)
record = pack_file.read(record_length)
except OSError as error:
raise ObjectCorruptError(obj_hash) from error
return self._decode_packed_record(obj_hash, record)
return None

def _read_index(
self, index_path: Path, pack_path: Path | None = None
) -> tuple[dict[str, tuple[int, int]], Path]:
try:
raw = index_path.read_bytes()
header_size = self._INDEX_HEADER.size
if len(raw) < header_size:
raise ValueError
magic, version, count = self._INDEX_HEADER.unpack(raw[:header_size])
if magic != self._INDEX_MAGIC or version != self._PACK_VERSION:
raise ValueError
expected_size = header_size + count * self._INDEX_ENTRY.size
if len(raw) != expected_size:
raise ValueError
entries: dict[str, tuple[int, int]] = {}
previous = ""
for position in range(count):
start = header_size + position * self._INDEX_ENTRY.size
raw_hash, offset, length = self._INDEX_ENTRY.unpack(
raw[start : start + self._INDEX_ENTRY.size]
)
obj_hash = raw_hash.hex()
if obj_hash <= previous or not length or obj_hash in entries:
raise ValueError
entries[obj_hash] = (offset, length)
previous = obj_hash
if pack_path is None:
pack_path = index_path.with_suffix(".pack")
self._validate_pack(pack_path, entries)
return entries, pack_path
except (OSError, ValueError, struct.error) as error:
raise ObjectCorruptError(index_path.name) from error

def _validate_pack(
self, pack_path: Path, entries: dict[str, tuple[int, int]] | None = None
) -> None:
try:
raw = pack_path.read_bytes()
if len(raw) < self._PACK_HEADER.size:
raise ValueError
magic, version, count = self._PACK_HEADER.unpack(raw[: self._PACK_HEADER.size])
if magic != self._PACK_MAGIC or version != self._PACK_VERSION:
raise ValueError
position = self._PACK_HEADER.size
records: list[tuple[int, int]] = []
for _ in range(count):
if position + self._RECORD_HEADER.size > len(raw):
raise ValueError
type_length, data_length, compressed_length = self._RECORD_HEADER.unpack(
raw[position : position + self._RECORD_HEADER.size]
)
record_length = self._RECORD_HEADER.size + type_length + compressed_length
if (
type_length not in {4, 5, 6}
or not data_length >= 0
or not compressed_length
or position + record_length > len(raw)
):
raise ValueError
records.append((position, record_length))
position += record_length
if position != len(raw):
raise ValueError
if entries is not None:
valid_records = set(records)
indexed_records = list(entries.values())
if any(entry not in valid_records for entry in indexed_records) or len(
set(indexed_records)
) != len(indexed_records):
raise ValueError
except (OSError, struct.error, ValueError) as error:
raise ObjectCorruptError(pack_path.name) from error

def _decode_packed_record(self, obj_hash: str, record: bytes) -> tuple[str, bytes]:
try:
if len(record) < self._RECORD_HEADER.size:
raise ValueError
type_length, data_length, compressed_length = self._RECORD_HEADER.unpack(
record[: self._RECORD_HEADER.size]
)
header_end = self._RECORD_HEADER.size + type_length
if len(record) != header_end + compressed_length:
raise ValueError
obj_type = record[self._RECORD_HEADER.size : header_end].decode("ascii")
if obj_type not in self._OBJECT_TYPES or data_length < 0:
raise ValueError
compressed = record[header_end:]
decompressor = zlib.decompressobj()
content = decompressor.decompress(compressed) + decompressor.flush()
if not decompressor.eof or decompressor.unused_data or decompressor.unconsumed_tail:
raise ValueError
if len(content) != data_length or self.hash_object(content, obj_type) != obj_hash:
raise ValueError
except (UnicodeDecodeError, ValueError, zlib.error) as error:
raise ObjectCorruptError(obj_hash) from error
return obj_type, content

def _object_path(self, hash: str) -> Path:
Expand Down
7 changes: 6 additions & 1 deletion tests/test_checkout_safety.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,12 @@ def test_checkout_rejects_leaf_symlink(tmp_path, dangling):
old = wt.store.write_object(b"original", "blob")
if not dangling:
wt.write_index([IndexEntry("100644", old, "link")])
(root / "link").symlink_to(outside)
try:
(root / "link").symlink_to(outside)
except OSError as error:
if getattr(error, "winerror", None) == 1314:
pytest.skip("Windows symlink privilege is unavailable")
raise
with pytest.raises(MiniGitError):
wt.checkout(target_tree(wt, "link"))
assert (root / "link").is_symlink()
Expand Down
3 changes: 2 additions & 1 deletion tests/test_commits.py
Original file line number Diff line number Diff line change
Expand Up @@ -387,7 +387,8 @@ def make_real_manager(temp_path):

def commit_file(m, path, contents, message):
"""Write `contents` to `path`, stage it, and commit on the current branch."""
(Path(m.root) / path).write_text(contents)
with open(Path(m.root) / path, "w", encoding="utf-8", newline="") as file:
file.write(contents)
m.tree.stage_file(path)
parent = m.read_ref(m._current_branch())
parents = [parent] if parent else []
Expand Down
Loading
Loading