diff --git a/README.md b/README.md index 6b9dd0d..9424b67 100644 --- a/README.md +++ b/README.md @@ -41,6 +41,39 @@ and find the virtual environment themselves, so activating it is optional. | `scripts/build.sh` | Builds an installable package into `dist/`. | | `scripts/install-git-hooks.sh` | Points Git at the hooks in `githooks/`. `init.sh` already does this. | +## Current workflow + +After setup, try this in a separate scratch directory with `minigit` on your PATH: + +```bash +minigit init +printf 'A\n' > file.txt +minigit add file.txt +minigit commit -m "A" +minigit branch feature +minigit checkout feature +printf 'B\n' > file.txt +minigit add file.txt +minigit commit -m "B" +minigit checkout main # file.txt contains A +minigit merge feature # fast-forward; file.txt contains B +minigit status # clean +``` + +Checkout restores files, executable modes, and the index. It rejects staged or +unstaged changes and untracked paths that would be overwritten. Fast-forward +merges keep HEAD on the current branch and create no additional commit. +Diverged branches report that three-way merging is not implemented yet. + +To test push, initialize another scratch repository and run +`minigit serve --port 9418 --token demo-token` there. From the first repository, +run `minigit push 127.0.0.1:9418 main --token demo-token`. Push transfers missing +objects and advances the remote branch only after validating its history; it +does not check out files or update the server's HEAD or index. Repeating the push +reports that the branch is up to date. Use a disposable token for this local demo. +Pull remains a handshake placeholder; fetching and restoring a branch through +`minigit pull` is not implemented yet. + ## Layout ``` diff --git a/minigit/remote.py b/minigit/remote.py index 0a04a00..566ef2f 100644 --- a/minigit/remote.py +++ b/minigit/remote.py @@ -65,6 +65,37 @@ def recv_exact(sock, buf: bytearray, size: int) -> bytes: return data +def _has_object(store, obj_hash: str) -> bool: + """Return whether `obj_hash` already exists in `store`. + + Prefers `store.has_object` (Module 1, #22) once the injected store has + grown one; falls back to a real read so this works against `main` today. + """ + + has_object = getattr(store, "has_object", None) + if has_object is not None: + return has_object(obj_hash) + try: + store.read_object(obj_hash) + return True + except ObjectNotFoundError: + return False + + +def _is_ancestor(commits, ancestor_hash: str, descendant_hash: str) -> bool: + """Return whether `ancestor_hash` is reachable from `descendant_hash`. + + Prefers `commits.is_ancestor` (Module 3, #24) once the injected commits + manager has grown one; falls back to `walk_history` so this works + against `main` today. + """ + + is_ancestor = getattr(commits, "is_ancestor", None) + if is_ancestor is not None: + return is_ancestor(ancestor_hash, descendant_hash) + return ancestor_hash in commits.walk_history(descendant_hash) + + class RemoteClient: """Push and pull commits between two minigit repos over a TCP connection.""" @@ -97,12 +128,25 @@ def _parse_address(self, address: str) -> tuple[str, int]: raise NetworkProtocolError(f"address must be host:port, got {address!r}") def push(self, remote_address: str, branch: str, token: str) -> None: - """Send local commits on `branch` to the remote, rejecting if it has diverged.""" + """Send local commits on `branch` to the remote, rejecting if it has diverged. + + Resolves the local tip and authenticates, then reads the remote's + current tip through `REF`. Equal tips mean nothing to do; a remote + tip that is not an ancestor of the local tip means someone else + pushed first, and the push is rejected. Otherwise every object + reachable from `branch` that the remote doesn't already have is + uploaded (`HAVE`/`PUT`), and the ref only moves after `DONE` reports + the transfer validated on the far side. + """ host, port = self._parse_address(remote_address) if len(token) == 0: raise NetworkProtocolError("push needs a token: pass --token") + local_hash = self.commits.read_ref(branch) + if local_hash is None: + raise NetworkProtocolError(f"local branch {branch!r} has no commits to push") + try: sock = socket.create_connection((host, port), timeout=5) except OSError as exc: @@ -117,16 +161,62 @@ def push(self, remote_address: str, branch: str, token: str) -> None: send_line(sock, f"REF {branch}") reply = receive_line(sock, buf) - remote_hash = reply.rsplit(" ", 1)[1] - print(f"remote {branch} is at {remote_hash}") - print("# Week 6 - send missing objects, move the ref last") + fields = reply.split(" ") + if len(fields) != 3 or fields[:2] != ["REF", branch]: + raise NetworkProtocolError(f"invalid REF response: {reply!r}") + remote_hash = fields[2] + if remote_hash != "-" and ( + len(remote_hash) != 40 or any(c not in "0123456789abcdef" for c in remote_hash) + ): + raise NetworkProtocolError(f"invalid remote hash: {remote_hash!r}") + + if remote_hash == local_hash: + print(f"{branch} is up to date") + return + + if remote_hash != "-" and not _is_ancestor(self.commits, remote_hash, local_hash): + raise NetworkProtocolError( + f"remote {branch} has diverged from local: " + f"{remote_hash} is not an ancestor of {local_hash}" + ) + + send_line(sock, f"PUSH {branch} {remote_hash} {local_hash}") + reply = receive_line(sock, buf) + if reply != "OK": + raise NetworkProtocolError(f"push rejected: {reply}") + + for obj_hash in self.collect_reachable(branch): + send_line(sock, f"HAVE {obj_hash}") + reply = receive_line(sock, buf) + if reply == "NO": + self._put_object(sock, buf, obj_hash) + elif reply != "YES": + raise NetworkProtocolError(f"expected YES or NO, got {reply!r}") + + send_line(sock, "DONE") + reply = receive_line(sock, buf) + if reply != "OK": + raise NetworkProtocolError(f"push failed: {reply}") + + print(f"{branch} now at {local_hash}") + except NetworkProtocolError: + raise + except (OSError, ObjectNotFoundError, ObjectCorruptError) as exc: + raise NetworkProtocolError(f"push to {host}:{port} failed: {exc}") from exc finally: sock.close() - # remote hash not an ancestor of local -> someone else pushed first -> NetworkProtocolError - # walk local commit graph from remote's hash up to local -> collect reachable objects - # send only the missing objects - # move the remote ref LAST, only after every object arrived + def _put_object(self, sock, buf: bytearray, obj_hash: str) -> None: + """Send `PUT ` followed by the object's `OBJ` header and bytes.""" + + obj_type, content = self.store.read_object(obj_hash) + send_line(sock, f"PUT {obj_hash}") + send_line(sock, f"OBJ {obj_type} {len(content)}") + sock.sendall(content) + + reply = receive_line(sock, buf) + if reply != "OK": + raise NetworkProtocolError(f"remote rejected {obj_hash}: {reply}") def pull(self, remote_address: str, branch: str, token: str) -> None: """Fetch `branch` from the remote and update the matching local ref.""" @@ -249,11 +339,14 @@ def _collect_tree(self, tree_hash: str, reachable: set[str]) -> None: class RemoteServer: """Accepts a RemoteClient's AUTH + REF handshake over TCP, one client at a time.""" - def __init__(self, repo_path=".", token="", host="127.0.0.1", port=0, store=None): + def __init__(self, repo_path=".", token="", host="127.0.0.1", port=0, store=None, commits=None): self.repo_path = repo_path self.token = token self.host = host self.store = store if store is not None else ObjectStore(repo_path) + self.commits = ( + commits if commits is not None else CommitManager(repo_path, store=self.store) + ) self._sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) self._sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) @@ -285,6 +378,11 @@ def serve_forever(self) -> None: try: self._handle_client(conn) + except (ObjectNotFoundError, ObjectCorruptError) as exc: + try: + send_line(conn, f"ERR object validation failed: {exc}") + except OSError: + pass except (NetworkProtocolError, OSError): pass # a bad client must not take down the server finally: @@ -312,8 +410,11 @@ def _handle_client(self, conn) -> None: self._send_ref(conn, value) elif command == "WANT": self._send_object(conn, value) + elif command == "PUSH": + self._handle_push(conn, buf, value) + return else: - send_line(conn, "ERR expected REF, WANT, or DONE") + send_line(conn, "ERR expected REF, WANT, PUSH, or DONE") return def _send_ref(self, conn, branch: str) -> None: @@ -326,14 +427,135 @@ def _send_ref(self, conn, branch: str) -> None: ): send_line(conn, "ERR invalid branch") return - ref_path = os.path.join(self.repo_path, ".minigit", "refs", "heads", branch) - if os.path.exists(ref_path): - with open(ref_path) as f: - commit_hash = f.read().strip() - else: - commit_hash = "-" + commit_hash = self.commits.read_ref(branch) or "-" send_line(conn, f"REF {branch} {commit_hash}") + def _handle_push(self, conn, buf: bytearray, args: str) -> None: + """Handle a PUSH sub-session: HAVE/PUT exchange, then validate and move the ref. + + `args` is `" "`. A failure anywhere in + here - a stale old ref, a bad object, a dropped connection - leaves + the ref untouched; any objects already stored are simply left in + place for a retry to reuse. + """ + + try: + branch, old_hash, new_hash = args.split(" ") + except ValueError: + send_line(conn, "ERR malformed PUSH") + return + + if ( + not branch + or any(part in {"", ".", ".."} for part in branch.split("/")) + or any(char in branch for char in "\0\\\r\n") + ): + send_line(conn, "ERR invalid branch") + return + if len(new_hash) != 40 or any(c not in "0123456789abcdef" for c in new_hash): + send_line(conn, "ERR invalid new hash") + return + + expected_old = None if old_hash == "-" else old_hash + if self.commits.read_ref(branch) != expected_old: + send_line(conn, f"ERR {branch} changed: expected {old_hash}") + return + + send_line(conn, "OK") + + while True: + line = receive_line(conn, buf) + command, _, value = line.partition(" ") + + if command == "HAVE": + send_line(conn, "YES" if _has_object(self.store, value) else "NO") + elif command == "PUT": + if not self._receive_pushed_object(conn, buf, value): + return + elif command == "DONE": + break + else: + send_line(conn, "ERR expected HAVE, PUT, or DONE") + return + + # Recheck the ref hasn't moved since we accepted the PUSH - another + # client's push could have landed while we were receiving objects. + if self.commits.read_ref(branch) != expected_old: + send_line(conn, "ERR ref changed during push") + return + + try: + self._validate_history(new_hash) + except (ObjectNotFoundError, ObjectCorruptError) as exc: + send_line(conn, f"ERR incomplete history: {exc}") + return + + if expected_old is not None and not _is_ancestor(self.commits, expected_old, new_hash): + send_line(conn, "ERR not a fast-forward") + return + + self.commits.write_ref(branch, new_hash) + send_line(conn, "OK") + + def _receive_pushed_object(self, conn, buf: bytearray, expected_hash: str) -> bool: + """Read one `PUT`'s `OBJ` header and bytes, hash-check, and store it. + + Replies `OK`/`ERR` itself (mirroring `_send_object`'s style) and + returns whether it succeeded, so the caller knows to stop the + session on failure. + """ + + if len(expected_hash) != 40 or any(c not in "0123456789abcdef" for c in expected_hash): + send_line(conn, "ERR invalid object hash") + return False + + header = receive_line(conn, buf) + command, _, rest = header.partition(" ") + if command != "OBJ": + send_line(conn, "ERR expected OBJ") + return False + + obj_type, _, length_text = rest.partition(" ") + if obj_type not in {"blob", "tree", "commit"} or not ( + length_text.isascii() and length_text.isdigit() + ): + send_line(conn, "ERR malformed OBJ header") + return False + + content = recv_exact(conn, buf, int(length_text)) + if self.store.hash_object(content, obj_type) != expected_hash: + send_line(conn, f"ERR hash mismatch for {expected_hash}") + return False + + self.store.write_object(content, obj_type) + send_line(conn, "OK") + return True + + def _validate_history(self, commit_hash: str) -> None: + """Verify every commit, tree, and blob reachable from `commit_hash` exists + with the expected type, before the ref is allowed to move. + """ + + visited_trees: set[str] = set() + for c_hash in self.commits.walk_history(commit_hash): + commit = self.commits.read_commit(c_hash) + self._validate_tree(commit.tree, visited_trees) + + def _validate_tree(self, tree_hash: str, visited: set[str]) -> None: + """Add `tree_hash` to `visited` and check everything nested under it, once each.""" + + if tree_hash in visited: + return + visited.add(tree_hash) + + for entry in self.store.read_tree(tree_hash): + if entry.type == "tree": + self._validate_tree(entry.hash, visited) + else: + obj_type, _ = self.store.read_object(entry.hash) + if obj_type != "blob": + raise ObjectCorruptError(entry.hash) + def _send_object(self, conn, obj_hash: str) -> None: """Reply with the requested object's bytes, or `ERR` if it isn't in the store.""" @@ -354,15 +576,22 @@ def _send_object(self, conn, obj_hash: str) -> None: conn.sendall(content) -# Wire protocol (draft only - Week 2 makes this real): -# One message per line, UTF-8 encoded, terminated with "\n". +# Wire protocol: one message per line, UTF-8 encoded, terminated with "\n". +# +# AUTH - client authenticates the connection with its token +# REF - ask for / report the commit hash a branch currently points to +# WANT - request the object with this hash (fetch/pull) +# PUSH - propose moving from ("-" = unborn) to +# HAVE - ask whether the remote already has an object +# PUT - announce an upload of , followed by an OBJ header + bytes +# OBJ - announces an object is coming next: its type and byte length +# DONE - no more messages from this side (also closes a PUSH session) +# ERR - something went wrong # -# AUTH - client authenticates the connection with its token -# REF - ask for / report the commit hash a branch currently points to -# WANT - request the object with this hash -# OBJ - announces an object is coming next: its type and byte length -# DONE - no more messages from this side -# ERR - something went wrong +# A push session: AUTH -> REF (read the remote's current tip) -> PUSH +# (propose the move) -> repeated HAVE/PUT (upload only what the remote is +# missing) -> DONE (remote validates the new history and moves the ref) -> +# connection closes. def register_subcommands(subparsers) -> None: diff --git a/tests/test_remote.py b/tests/test_remote.py index 3ab7652..97be842 100644 --- a/tests/test_remote.py +++ b/tests/test_remote.py @@ -1,12 +1,16 @@ import hashlib import socket import threading +from dataclasses import dataclass +from pathlib import Path from typing import NamedTuple import pytest +from minigit.commits import CommitManager from minigit.errors import NetworkProtocolError from minigit.objects import ObjectStore +from minigit.objects import TreeEntry as RealTreeEntry from minigit.remote import RemoteClient, RemoteServer, receive_line, recv_exact, send_line KNOWN_HASH = "a" * 40 @@ -90,36 +94,275 @@ def test_push_empty_token_raises(): client.push("127.0.0.1:9418", "main", "") -def test_push_correct_token_does_not_raise(remote_server): - client = make_client() - client.push(f"127.0.0.1:{remote_server.port}", "main", "tok") +# --- push: real M1/M3 instances on both ends, exercising the full PUSH/HAVE/PUT/DONE session --- -def test_push_wrong_token_raises(remote_server): - client = make_client() +@dataclass +class PushPair: + client: RemoteClient + local_store: ObjectStore + local_commits: CommitManager + remote_store: ObjectStore + remote_commits: CommitManager + server: RemoteServer + address: str + + +@pytest.fixture +def push_pair(tmp_path): + remote_root = tmp_path / "remote" + remote_root.mkdir() + remote_store = ObjectStore(str(remote_root)) + remote_commits = CommitManager(str(remote_root), store=remote_store) + server = RemoteServer( + repo_path=str(remote_root), token="tok", port=0, store=remote_store, commits=remote_commits + ) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + + local_root = tmp_path / "local" + local_root.mkdir() + local_store = ObjectStore(str(local_root)) + local_commits = CommitManager(str(local_root), store=local_store) + client = RemoteClient(repo_path=str(local_root), store=local_store, commits=local_commits) + + yield PushPair( + client=client, + local_store=local_store, + local_commits=local_commits, + remote_store=remote_store, + remote_commits=remote_commits, + server=server, + address=f"127.0.0.1:{server.port}", + ) + + server.close() + thread.join(timeout=2) + + +def _write_tree(store: ObjectStore, files: dict) -> str: + """Build (possibly nested) tree objects for a flat {path: content} mapping.""" + + entries = [] + subdirs: dict[str, dict] = {} + for path, content in files.items(): + top, sep, rest = path.partition("/") + if sep: + subdirs.setdefault(top, {})[rest] = content + else: + blob_hash = store.write_object(content, "blob") + entries.append(RealTreeEntry("100644", "blob", blob_hash, top)) + for name, nested in subdirs.items(): + entries.append(RealTreeEntry("40000", "tree", _write_tree(store, nested), name)) + return store.write_tree(entries) + + +def _commit( + store: ObjectStore, commits: CommitManager, files: dict, parents: list, message="msg" +) -> str: + tree_hash = _write_tree(store, files) + return commits.create_commit(tree_hash, parents, "tester ", message) + + +def test_push_wrong_token_raises(push_pair): + _commit(push_pair.local_store, push_pair.local_commits, {"a.txt": b"1"}, []) with pytest.raises(NetworkProtocolError): - client.push(f"127.0.0.1:{remote_server.port}", "main", "wrong") + push_pair.client.push(push_pair.address, "main", "wrong") + assert push_pair.remote_commits.read_ref("main") is None -def test_push_existing_branch_prints_matching_hash(remote_server, capsys): - client = make_client() - client.push(f"127.0.0.1:{remote_server.port}", "main", "tok") - assert KNOWN_HASH in capsys.readouterr().out +def test_push_missing_local_branch_raises(push_pair): + with pytest.raises(NetworkProtocolError): + push_pair.client.push(push_pair.address, "nope", "tok") -def test_push_missing_branch_prints_dash(remote_server, capsys): - client = make_client() - client.push(f"127.0.0.1:{remote_server.port}", "nope", "tok") - assert "is at -" in capsys.readouterr().out +def test_push_closed_port_raises_network_protocol_error(push_pair): + _commit(push_pair.local_store, push_pair.local_commits, {"a.txt": b"1"}, []) + push_pair.server.close() # closing the listening socket refuses new connections immediately + with pytest.raises(NetworkProtocolError): + push_pair.client.push(push_pair.address, "main", "tok") -def test_push_closed_port_raises_network_protocol_error(remote_server): - port = remote_server.port - remote_server.close() - remote_server._thread.join(timeout=2) # wait for serve_forever() to actually exit - client = make_client() +def test_push_nested_files_multiple_commits(push_pair): + pair = push_pair + c1 = _commit(pair.local_store, pair.local_commits, {"a.txt": b"one"}, []) + c2 = _commit( + pair.local_store, + pair.local_commits, + {"a.txt": b"one", "src/main.py": b"print(1)"}, + [c1], + ) + + pair.client.push(pair.address, "main", "tok") + + assert pair.remote_commits.read_ref("main") == c2 + tree_entries = pair.remote_store.read_tree(pair.remote_commits.read_commit(c2).tree) + src_entry = next(e for e in tree_entries if e.name == "src") + nested = pair.remote_store.read_tree(src_entry.hash) + assert nested[0].name == "main.py" + + # server never touches its own HEAD/index while handling a push + assert not (Path(pair.server.repo_path) / ".minigit" / "HEAD").exists() + assert not (Path(pair.server.repo_path) / ".minigit" / "index").exists() + + +def test_push_to_empty_branch_then_one_more_commit(push_pair): + pair = push_pair + assert pair.remote_commits.read_ref("main") is None + + c1 = _commit(pair.local_store, pair.local_commits, {"a.txt": b"one"}, []) + pair.client.push(pair.address, "main", "tok") + assert pair.remote_commits.read_ref("main") == c1 + + c2 = _commit(pair.local_store, pair.local_commits, {"a.txt": b"two"}, [c1]) + pair.client.push(pair.address, "main", "tok") + assert pair.remote_commits.read_ref("main") == c2 + + +def test_push_up_to_date_reports_and_stops(push_pair, capsys): + pair = push_pair + _commit(pair.local_store, pair.local_commits, {"a.txt": b"one"}, []) + pair.client.push(pair.address, "main", "tok") + capsys.readouterr() + + pair.client.push(pair.address, "main", "tok") + assert "up to date" in capsys.readouterr().out + + +def test_repeat_push_transfers_no_objects(push_pair): + pair = push_pair + _commit(pair.local_store, pair.local_commits, {"a.txt": b"one"}, []) + pair.client.push(pair.address, "main", "tok") + + put_calls = [] + original_put = RemoteClient._put_object + RemoteClient._put_object = lambda self, sock, buf, obj_hash: put_calls.append(obj_hash) + try: + pair.client.push(pair.address, "main", "tok") + finally: + RemoteClient._put_object = original_put + + assert put_calls == [] + + +def test_shared_blobs_transfer_once_and_receiver_reads_full_history(push_pair): + pair = push_pair + c1 = _commit(pair.local_store, pair.local_commits, {"shared.txt": b"same", "a.txt": b"1"}, []) + c2 = _commit(pair.local_store, pair.local_commits, {"shared.txt": b"same", "a.txt": b"2"}, [c1]) + + put_calls = [] + original_put = RemoteClient._put_object + + def counting_put(self, sock, buf, obj_hash): + put_calls.append(obj_hash) + return original_put(self, sock, buf, obj_hash) + + RemoteClient._put_object = counting_put + try: + pair.client.push(pair.address, "main", "tok") + finally: + RemoteClient._put_object = original_put + + assert len(put_calls) == len(set(put_calls)) # every object uploaded at most once + + # a *fresh* pair of M1/M3 instances, as if a different process opened the repo + fresh_store = ObjectStore(pair.server.repo_path) + fresh_commits = CommitManager(pair.server.repo_path, store=fresh_store) + assert fresh_commits.read_ref("main") == c2 + assert set(fresh_commits.walk_history(c2)) == {c1, c2} + shared_hash = pair.local_store.hash_object(b"same", "blob") + assert fresh_store.read_object(shared_hash) == ("blob", b"same") + + +def test_push_bad_auth_leaves_ref_unchanged(push_pair): + pair = push_pair + _commit(pair.local_store, pair.local_commits, {"a.txt": b"1"}, []) with pytest.raises(NetworkProtocolError): - client.push(f"127.0.0.1:{port}", "main", "tok") + pair.client.push(pair.address, "main", "wrong-token") + assert pair.remote_commits.read_ref("main") is None + + +def test_push_bad_object_payload_leaves_ref_unchanged(push_pair): + pair = push_pair + c1 = _commit(pair.local_store, pair.local_commits, {"a.txt": b"1"}, []) + + sock = socket.create_connection(("127.0.0.1", pair.server.port), timeout=5) + buf = bytearray() + try: + send_line(sock, "AUTH tok") + assert receive_line(sock, buf) == "OK" + send_line(sock, f"PUSH main - {c1}") + assert receive_line(sock, buf) == "OK" + send_line(sock, f"HAVE {c1}") + assert receive_line(sock, buf) == "NO" + send_line(sock, f"PUT {c1}") + bogus = b"not the right content" + send_line(sock, f"OBJ commit {len(bogus)}") + sock.sendall(bogus) + assert receive_line(sock, buf).startswith("ERR") + finally: + sock.close() + + assert pair.remote_commits.read_ref("main") is None + + +def test_push_disconnect_before_done_leaves_ref_unchanged(push_pair): + pair = push_pair + c1 = _commit(pair.local_store, pair.local_commits, {"a.txt": b"1"}, []) + + sock = socket.create_connection(("127.0.0.1", pair.server.port), timeout=5) + buf = bytearray() + send_line(sock, "AUTH tok") + assert receive_line(sock, buf) == "OK" + send_line(sock, f"PUSH main - {c1}") + assert receive_line(sock, buf) == "OK" + send_line(sock, f"HAVE {c1}") + assert receive_line(sock, buf) == "NO" + sock.close() # disconnect instead of PUT + DONE + + assert pair.remote_commits.read_ref("main") is None + + +def test_push_rejects_unrelated_remote_tip(push_pair): + pair = push_pair + unrelated = _commit(pair.remote_store, pair.remote_commits, {"other.txt": b"x"}, []) + _commit(pair.local_store, pair.local_commits, {"a.txt": b"1"}, []) + + with pytest.raises(NetworkProtocolError): + pair.client.push(pair.address, "main", "tok") + assert pair.remote_commits.read_ref("main") == unrelated + + +def test_push_rejects_ref_changed_during_transfer(push_pair): + pair = push_pair + c1 = _commit(pair.local_store, pair.local_commits, {"a.txt": b"1"}, []) + + sock = socket.create_connection(("127.0.0.1", pair.server.port), timeout=5) + buf = bytearray() + try: + send_line(sock, "AUTH tok") + assert receive_line(sock, buf) == "OK" + send_line(sock, f"PUSH main - {c1}") + assert receive_line(sock, buf) == "OK" + + # simulate another push landing on the remote mid-transfer + pair.remote_commits.write_ref("main", "b" * 40) + + for obj_hash in pair.client.collect_reachable("main"): + send_line(sock, f"HAVE {obj_hash}") + reply = receive_line(sock, buf) + if reply == "NO": + obj_type, content = pair.local_store.read_object(obj_hash) + send_line(sock, f"PUT {obj_hash}") + send_line(sock, f"OBJ {obj_type} {len(content)}") + sock.sendall(content) + assert receive_line(sock, buf) == "OK" + send_line(sock, "DONE") + assert receive_line(sock, buf).startswith("ERR") + finally: + sock.close() + + assert pair.remote_commits.read_ref("main") == "b" * 40 def test_serve_objects_returns_requested_object(remote_server, tmp_path): @@ -446,3 +689,48 @@ def test_server_survives_corrupt_objects_and_invalid_requests(remote_server): assert receive_line(sock, buf) == "OBJ blob 4" assert recv_exact(sock, buf, 4) == b"good" send_line(sock, "DONE") + + +@pytest.mark.parametrize("command", ["HAVE", "PUT"]) +def test_push_corrupt_existing_object_does_not_stop_server(push_pair, command): + pair = push_pair + tip = _commit(pair.local_store, pair.local_commits, {"a.txt": b"one"}, []) + blob = pair.remote_store.write_object(b"one", "blob") + pair.remote_store._object_path(blob).write_bytes(b"corrupt") + with socket.create_connection(("127.0.0.1", pair.server.port), timeout=2) as sock: + buf = bytearray() + send_line(sock, "AUTH tok") + assert receive_line(sock, buf) == "OK" + send_line(sock, f"PUSH main - {tip}") + assert receive_line(sock, buf) == "OK" + send_line(sock, f"{command} {blob}") + if command == "PUT": + send_line(sock, "OBJ blob 3") + sock.sendall(b"one") + assert receive_line(sock, buf).startswith("ERR") + assert pair.remote_commits.read_ref("main") is None + with socket.create_connection(("127.0.0.1", pair.server.port), timeout=2) as sock: + buf = bytearray() + send_line(sock, "AUTH tok") + assert receive_line(sock, buf) == "OK" + send_line(sock, "REF main") + assert receive_line(sock, buf) == "REF main -" + send_line(sock, "DONE") + + +@pytest.mark.parametrize("reply", ["ERR", "ERR invalid branch", "REF other -", "REF main bad"]) +def test_push_invalid_ref_reply_is_protocol_error(tmp_path, monkeypatch, reply): + manager = CommitManager(tmp_path) + _commit(manager.store, manager, {"a.txt": b"one"}, []) + + class ReplySocket(FakeReceiveSocket): + def sendall(self, data): + pass + + def close(self): + pass + + sock = ReplySocket(f"OK\n{reply}\n".encode()) + monkeypatch.setattr(socket, "create_connection", lambda *args, **kwargs: sock) + with pytest.raises(NetworkProtocolError, match="invalid"): + RemoteClient(tmp_path).push("localhost:9418", "main", "tok") diff --git a/tests/test_week4_integration.py b/tests/test_week4_integration.py new file mode 100644 index 0000000..faf7810 --- /dev/null +++ b/tests/test_week4_integration.py @@ -0,0 +1,54 @@ +import threading +from pathlib import Path + +from minigit.cli import main +from minigit.commits import CommitManager +from minigit.remote import RemoteClient, RemoteServer + + +def test_week4_cli_checkpoint_and_push(tmp_path, monkeypatch, capsys): + local = tmp_path / "local" + local.mkdir() + monkeypatch.chdir(local) + assert main(["init"]) == 0 + (local / "file.txt").write_text("A\n") + assert main(["add", "file.txt"]) == 0 + assert main(["commit", "-m", "A"]) == 0 + manager = CommitManager(local) + a = manager.read_ref("main") + assert main(["branch", "feature"]) == 0 + assert main(["checkout", "feature"]) == 0 + (local / "file.txt").write_text("B\n") + assert main(["add", "file.txt"]) == 0 + assert main(["commit", "-m", "B"]) == 0 + b = manager.read_ref("feature") + assert a != b + assert main(["checkout", "main"]) == 0 + assert (local / "file.txt").read_text() == "A\n" + assert main(["merge", "feature"]) == 0 + assert (local / "file.txt").read_text() == "B\n" + assert manager.read_head() == "main" + assert manager.read_ref("main") == manager.read_ref("feature") == b + capsys.readouterr() + assert main(["status"]) == 0 + assert capsys.readouterr().out.strip() == "clean" + + remote = tmp_path / "remote" + remote.mkdir() + server = RemoteServer(remote, token="test-token") + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + assert main(["push", f"127.0.0.1:{server.port}", "main", "--token", "test-token"]) == 0 + receiver = CommitManager(remote) + assert receiver.read_ref("main") == b + assert receiver.walk_history(b) == [b, a] + for obj_hash in RemoteClient(local).collect_reachable("main"): + assert receiver.store.read_object(obj_hash) == manager.store.read_object(obj_hash) + assert not (remote / ".minigit/index").exists() + assert not (remote / ".minigit/HEAD").exists() + assert list(Path(remote).iterdir()) == [remote / ".minigit"] + finally: + server.close() + thread.join(timeout=2) + assert not thread.is_alive()