diff --git a/minigit/remote.py b/minigit/remote.py index ad80ff6..0a04a00 100644 --- a/minigit/remote.py +++ b/minigit/remote.py @@ -9,12 +9,12 @@ Build the `RemoteClient` class here, per the interface contract. """ -# from .objects import ObjectStore (Module 1 & 3) -# from .commits import CommitManager import os import socket -from minigit.errors import NetworkProtocolError +from minigit.commits import CommitManager +from minigit.errors import NetworkProtocolError, ObjectCorruptError, ObjectNotFoundError +from minigit.objects import ObjectStore def send_line(sock, text: str) -> None: @@ -40,21 +40,41 @@ def receive_line(sock, buf: bytearray) -> str: line, _, rest = buf.partition(b"\n") buf[:] = rest - return line.decode() + try: + return line.decode("utf-8") + except UnicodeDecodeError as exc: + raise NetworkProtocolError("protocol line is not UTF-8") from exc + + +def recv_exact(sock, buf: bytearray, size: int) -> bytes: + """Read exactly `size` bytes from `sock`, sharing `buf` with `receive_line`. + + Same buffer contract: bytes past the `size`th belong to whatever message + comes next on this connection and are left in `buf` for that call to + consume, instead of being read (and discarded) here. + """ + + while len(buf) < size: + chunk = sock.recv(4096) + if not chunk: + raise NetworkProtocolError("connection closed mid-message") + buf.extend(chunk) + + data = bytes(buf[:size]) + del buf[:size] + return data class RemoteClient: """Push and pull commits between two minigit repos over a TCP connection.""" def __init__(self, repo_path=".", store=None, commits=None): - self.repo_path = repo_path self.config_path = os.path.join(self.repo_path, ".minigit", "config") - self.store = store - self.commits = commits - # Below is correct but need name of function within module 1 & 3 - # self.store = store if store is not None else ObjectStore(self.repo_path) - # self.commits = commits if commits is not None else CommitManager(self.repo_path) + self.store = store if store is not None else ObjectStore(self.repo_path) + self.commits = ( + commits if commits is not None else CommitManager(self.repo_path, store=self.store) + ) def _parse_address(self, address: str) -> tuple[str, int]: """split a string address by host part(string) and the port part(integer)""" @@ -135,14 +155,105 @@ def pull(self, remote_address: str, branch: str, token: str) -> None: finally: sock.close() + def fetch_objects(self, remote_address: str, hashes: list[str], token: str) -> None: + """Fetch each of `hashes` from the remote and write it into the local store. + + Authenticates once, then sends one `WANT` per unique hash over the + same connection. Each reply's content is re-hashed and checked + against the hash that was requested before it is stored, so a + corrupted or mismatched reply never lands in the object store. + """ + + host, port = self._parse_address(remote_address) + if len(token) == 0: + raise NetworkProtocolError("fetch needs a token: pass --token") + + try: + sock = socket.create_connection((host, port), timeout=5) + except OSError as exc: + raise NetworkProtocolError(f"could not connect to {host}:{port}: {exc}") from exc + + try: + buf = bytearray() + send_line(sock, f"AUTH {token}") + reply = receive_line(sock, buf) + if reply != "OK": + raise NetworkProtocolError(f"auth failed: {reply}") + + for obj_hash in dict.fromkeys(hashes): + send_line(sock, f"WANT {obj_hash}") + self._receive_object(sock, buf, obj_hash) + + send_line(sock, "DONE") + except OSError as exc: + raise NetworkProtocolError(f"connection to {host}:{port} failed: {exc}") from exc + finally: + sock.close() + + def _receive_object(self, sock, buf: bytearray, expected_hash: str) -> None: + """Read one `OBJ`/`ERR` reply for `expected_hash` and store it if it checks out.""" + + header = receive_line(sock, buf) + command, _, rest = header.partition(" ") + + if command == "ERR": + raise NetworkProtocolError(f"remote could not provide {expected_hash}: {rest}") + if command != "OBJ": + raise NetworkProtocolError(f"expected OBJ, got {header!r}") + + obj_type, _, length_text = rest.partition(" ") + if obj_type not in {"blob", "tree", "commit"} or not ( + length_text.isascii() and length_text.isdigit() + ): + raise NetworkProtocolError(f"malformed OBJ header: {header!r}") + + content = recv_exact(sock, buf, int(length_text)) + if self.store.hash_object(content, obj_type) != expected_hash: + raise NetworkProtocolError(f"object {expected_hash} failed hash verification") + + self.store.write_object(content, obj_type) + + def collect_reachable(self, branch: str) -> set[str]: + """Return every commit, tree, and blob hash reachable from `branch`'s tip. + + Local-only this week: this is what push will later diff against the + remote's advertised hash to find what's actually missing there. + """ + + tip = self.commits.read_ref(branch) + if tip is None: + return set() + + reachable: set[str] = set() + for commit_hash in self.commits.walk_history(tip): + reachable.add(commit_hash) + commit = self.commits.read_commit(commit_hash) + self._collect_tree(commit.tree, reachable) + + return reachable + + def _collect_tree(self, tree_hash: str, reachable: set[str]) -> None: + """Add `tree_hash` and everything nested under it to `reachable`, once each.""" + + if tree_hash in reachable: + return + reachable.add(tree_hash) + + for entry in self.store.read_tree(tree_hash): + if entry.type == "tree": + self._collect_tree(entry.hash, reachable) + else: + reachable.add(entry.hash) + 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): + def __init__(self, repo_path=".", token="", host="127.0.0.1", port=0, store=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._sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) self._sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) @@ -180,7 +291,7 @@ def serve_forever(self) -> None: conn.close() def _handle_client(self, conn) -> None: - """Run one client's AUTH + REF handshake.""" + """Authenticate the connection, then answer REF / WANT requests until DONE.""" buf = bytearray() @@ -191,12 +302,30 @@ def _handle_client(self, conn) -> None: return send_line(conn, "OK") - line = receive_line(conn, buf) - command, _, branch = line.partition(" ") - if command != "REF": - send_line(conn, "ERR expected REF") - return + while True: + line = receive_line(conn, buf) + command, _, value = line.partition(" ") + if command == "DONE": + return + elif command == "REF": + self._send_ref(conn, value) + elif command == "WANT": + self._send_object(conn, value) + else: + send_line(conn, "ERR expected REF, WANT, or DONE") + return + + def _send_ref(self, conn, branch: str) -> None: + """Reply with the commit hash `branch` currently points at, or `-` if it has none.""" + + 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 ref_path = os.path.join(self.repo_path, ".minigit", "refs", "heads", branch) if os.path.exists(ref_path): with open(ref_path) as f: @@ -205,6 +334,25 @@ def _handle_client(self, conn) -> None: commit_hash = "-" send_line(conn, f"REF {branch} {commit_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.""" + + if len(obj_hash) != 40 or any(c not in "0123456789abcdef" for c in obj_hash): + send_line(conn, "ERR invalid object hash") + return + try: + obj_type, content = self.store.read_object(obj_hash) + except ObjectNotFoundError: + send_line(conn, f"ERR unknown object {obj_hash}") + return + + except ObjectCorruptError: + send_line(conn, f"ERR corrupt object {obj_hash}") + return + + send_line(conn, f"OBJ {obj_type} {len(content)}") + conn.sendall(content) + # Wire protocol (draft only - Week 2 makes this real): # One message per line, UTF-8 encoded, terminated with "\n". diff --git a/tests/test_remote.py b/tests/test_remote.py index 227aa98..3ab7652 100644 --- a/tests/test_remote.py +++ b/tests/test_remote.py @@ -1,15 +1,33 @@ +import hashlib +import socket import threading +from typing import NamedTuple import pytest from minigit.errors import NetworkProtocolError -from minigit.remote import RemoteClient, RemoteServer, receive_line +from minigit.objects import ObjectStore +from minigit.remote import RemoteClient, RemoteServer, receive_line, recv_exact, send_line KNOWN_HASH = "a" * 40 class FakeObjectStore: - pass + """Matches ObjectStore's hashing so wire-level tests can check hashes without disk I/O.""" + + def __init__(self): + self._objects = {} + + def hash_object(self, data: bytes, obj_type: str) -> str: + return hashlib.sha1(f"{obj_type} {len(data)}".encode() + b"\0" + data).hexdigest() + + def write_object(self, data: bytes, obj_type: str) -> str: + obj_hash = self.hash_object(data, obj_type) + self._objects[obj_hash] = (obj_type, data) + return obj_hash + + def read_object(self, hash: str) -> tuple[str, bytes]: + return self._objects[hash] class FakeCommitManager: @@ -104,6 +122,113 @@ def test_push_closed_port_raises_network_protocol_error(remote_server): client.push(f"127.0.0.1:{port}", "main", "tok") +def test_serve_objects_returns_requested_object(remote_server, tmp_path): + store = ObjectStore(str(tmp_path)) + blob_hash = store.write_object(b"hello world", "blob") + + sock = socket.create_connection(("127.0.0.1", remote_server.port), timeout=5) + buf = bytearray() + try: + send_line(sock, "AUTH tok") + assert receive_line(sock, buf) == "OK" + send_line(sock, "REF main") + receive_line(sock, buf) # REF main : covered by the push/pull tests + + send_line(sock, f"WANT {blob_hash}") + header = receive_line(sock, buf) + obj_type, length = header.removeprefix("OBJ ").split(" ") + assert obj_type == "blob" + assert recv_exact(sock, buf, int(length)) == b"hello world" + + send_line(sock, "DONE") + finally: + sock.close() + + +def test_serve_objects_unknown_hash_sends_err_and_stays_open(remote_server): + sock = socket.create_connection(("127.0.0.1", remote_server.port), timeout=5) + buf = bytearray() + try: + send_line(sock, "AUTH tok") + receive_line(sock, buf) + send_line(sock, "REF main") + receive_line(sock, buf) + + send_line(sock, "WANT " + "f" * 40) + assert receive_line(sock, buf).startswith("ERR") + + send_line(sock, "DONE") + finally: + sock.close() + + +def test_fetch_objects_stores_multiple_objects_from_one_connection(remote_server, tmp_path): + remote_store = ObjectStore(remote_server.repo_path) + blob_hash = remote_store.write_object(b"binary:\x00\nbytes", "blob") + commit_hash = remote_store.write_object(b"tree deadbeef", "commit") + + local_store = ObjectStore(str(tmp_path / "local")) + client = RemoteClient(store=local_store, commits=FakeCommitManager()) + + client.fetch_objects( + f"127.0.0.1:{remote_server.port}", [blob_hash, commit_hash, blob_hash], "tok" + ) + + assert local_store.read_object(blob_hash) == ("blob", b"binary:\x00\nbytes") + assert local_store.read_object(commit_hash) == ("commit", b"tree deadbeef") + + +def test_fetch_objects_unknown_hash_raises(remote_server, tmp_path): + client = RemoteClient(store=ObjectStore(str(tmp_path / "local")), commits=FakeCommitManager()) + with pytest.raises(NetworkProtocolError): + client.fetch_objects(f"127.0.0.1:{remote_server.port}", ["f" * 40], "tok") + + +def test_fetch_objects_empty_token_raises(): + client = make_client() + with pytest.raises(NetworkProtocolError): + client.fetch_objects("127.0.0.1:9418", ["a" * 40], "") + + +class FakeReceiveSocket: + """Hands back one pre-scripted OBJ (or ERR) reply for `_receive_object`.""" + + def __init__(self, reply: bytes): + self._reply = reply + + def recv(self, size): + chunk, self._reply = self._reply[:size], self._reply[size:] + return chunk + + +def test_receive_object_wrong_hash_raises(): + client = make_client() + sock = FakeReceiveSocket(b"OBJ blob 5\nnope!") + with pytest.raises(NetworkProtocolError): + client._receive_object(sock, bytearray(), "a" * 40) + + +def test_receive_object_err_reply_raises(): + client = make_client() + sock = FakeReceiveSocket(b"ERR unknown object\n") + with pytest.raises(NetworkProtocolError): + client._receive_object(sock, bytearray(), "a" * 40) + + +def test_receive_object_malformed_header_raises(): + client = make_client() + sock = FakeReceiveSocket(b"OBJ blob notanumber\n") + with pytest.raises(NetworkProtocolError): + client._receive_object(sock, bytearray(), "a" * 40) + + +def test_receive_object_dropped_connection_raises(): + client = make_client() + sock = FakeReceiveSocket(b"OBJ blob 20\nshort") + with pytest.raises(NetworkProtocolError): + client._receive_object(sock, bytearray(), "a" * 40) + + class FakeSocket: """Hands back pre-scripted chunks instead of reading a real socket.""" @@ -120,3 +245,204 @@ def test_receive_line_splits_two_lines_from_one_packet(): assert receive_line(sock, buf) == "AUTH tok" assert receive_line(sock, buf) == "REF main" + + +def test_recv_exact_reads_payload_split_across_packets(): + sock = FakeSocket([b"hel", b"lo!"]) + buf = bytearray() + + assert recv_exact(sock, buf, 6) == b"hello!" + + +def test_recv_exact_leaves_trailing_bytes_for_next_read(): + sock = FakeSocket([b"helloREF main\n"]) + buf = bytearray() + + assert recv_exact(sock, buf, 5) == b"hello" + assert receive_line(sock, buf) == "REF main" + + +def test_recv_exact_raises_on_connection_closed_mid_payload(): + sock = FakeSocket([b"he"]) + buf = bytearray() + + with pytest.raises(NetworkProtocolError): + recv_exact(sock, buf, 5) + + +# --- collect_reachable: fakes for the not-yet-merged M1 #16 / M3 #18 interfaces --- + + +class TreeEntry(NamedTuple): + mode: str + type: str + hash: str + name: str + + +class CommitData(NamedTuple): + tree: str + parents: list[str] + author: str + committer: str + message: str + + +class FakeTreeStore: + """`read_tree` double: hash -> entries, wired up directly by each test.""" + + def __init__(self): + self.trees: dict[str, list[TreeEntry]] = {} + self.read_tree_calls: list[str] = [] + + def read_tree(self, tree_hash: str) -> list[TreeEntry]: + self.read_tree_calls.append(tree_hash) + return self.trees[tree_hash] + + +class FakeCommitGraph: + """`read_ref`/`read_commit`/`walk_history` double, wired up directly by each test.""" + + def __init__(self): + self.refs: dict[str, str] = {} + self.commits: dict[str, CommitData] = {} + + def read_ref(self, branch: str) -> str | None: + return self.refs.get(branch) + + def read_commit(self, commit_hash: str) -> CommitData: + return self.commits[commit_hash] + + def walk_history(self, start_hash: str) -> list[str]: + order: list[str] = [] + seen: set[str] = set() + + def visit(commit_hash: str) -> None: + if commit_hash in seen: + return + seen.add(commit_hash) + order.append(commit_hash) + for parent in self.commits[commit_hash].parents: + visit(parent) + + visit(start_hash) + return order + + +def test_collect_reachable_unborn_branch_returns_empty_set(): + client = RemoteClient(store=FakeTreeStore(), commits=FakeCommitGraph()) + assert client.collect_reachable("main") == set() + + +def test_collect_reachable_single_commit_collects_tree_and_blobs(): + store = FakeTreeStore() + commits = FakeCommitGraph() + + store.trees["tree1"] = [ + TreeEntry("100644", "blob", "blobA", "a.txt"), + TreeEntry("100644", "blob", "blobB", "b.txt"), + ] + commits.commits["c1"] = CommitData("tree1", [], "a", "a", "root") + commits.refs["main"] = "c1" + + client = RemoteClient(store=store, commits=commits) + assert client.collect_reachable("main") == {"c1", "tree1", "blobA", "blobB"} + + +def test_collect_reachable_nested_trees(): + store = FakeTreeStore() + commits = FakeCommitGraph() + + store.trees["root-tree"] = [ + TreeEntry("100644", "blob", "readme", "README.md"), + TreeEntry("40000", "tree", "src-tree", "src"), + ] + store.trees["src-tree"] = [TreeEntry("100644", "blob", "mainpy", "main.py")] + commits.commits["c1"] = CommitData("root-tree", [], "a", "a", "root") + commits.refs["main"] = "c1" + + client = RemoteClient(store=store, commits=commits) + assert client.collect_reachable("main") == { + "c1", + "root-tree", + "readme", + "src-tree", + "mainpy", + } + + +def test_collect_reachable_dedupes_shared_blob_and_tree_across_commits(): + store = FakeTreeStore() + commits = FakeCommitGraph() + + store.trees["shared-tree"] = [TreeEntry("100644", "blob", "shared-blob", "f.txt")] + commits.commits["c1"] = CommitData("shared-tree", [], "a", "a", "root") + commits.commits["c2"] = CommitData("shared-tree", ["c1"], "a", "a", "unchanged") + commits.refs["main"] = "c2" + + client = RemoteClient(store=store, commits=commits) + reachable = client.collect_reachable("main") + + assert reachable == {"c1", "c2", "shared-tree", "shared-blob"} + assert store.read_tree_calls == ["shared-tree"] # walked once, not once per commit + + +def test_collect_reachable_includes_both_merge_parents(): + store = FakeTreeStore() + commits = FakeCommitGraph() + + store.trees["tree-a"] = [TreeEntry("100644", "blob", "blob-a", "a.txt")] + store.trees["tree-b"] = [TreeEntry("100644", "blob", "blob-b", "b.txt")] + store.trees["tree-root"] = [] + store.trees["tree-merge"] = [TreeEntry("100644", "blob", "blob-a", "a.txt")] + + commits.commits["root"] = CommitData("tree-root", [], "a", "a", "root") + commits.commits["left"] = CommitData("tree-a", ["root"], "a", "a", "left") + commits.commits["right"] = CommitData("tree-b", ["root"], "a", "a", "right") + commits.commits["merge"] = CommitData("tree-merge", ["left", "right"], "a", "a", "merge") + commits.refs["main"] = "merge" + + client = RemoteClient(store=store, commits=commits) + reachable = client.collect_reachable("main") + + assert reachable == { + "root", + "left", + "right", + "merge", + "tree-root", + "tree-a", + "tree-b", + "tree-merge", + "blob-a", + "blob-b", + } + + +def test_receive_line_invalid_utf8_is_protocol_error(): + with pytest.raises(NetworkProtocolError): + receive_line(FakeSocket([b"\xff\n"]), bytearray()) + + +@pytest.mark.parametrize("header", [b"OBJ blob \xc2\xb2\n", b"OBJ unknown 0\n"]) +def test_receive_object_rejects_invalid_type_and_unicode_length(header): + with pytest.raises(NetworkProtocolError): + make_client()._receive_object(FakeReceiveSocket(header), bytearray(), "a" * 40) + + +def test_server_survives_corrupt_objects_and_invalid_requests(remote_server): + store = ObjectStore(remote_server.repo_path) + bad_hash = store.write_object(b"broken", "blob") + store._object_path(bad_hash).write_bytes(b"not compressed") + good_hash = store.write_object(b"good", "blob") + with socket.create_connection(("127.0.0.1", remote_server.port), timeout=5) as sock: + buf = bytearray() + send_line(sock, "AUTH tok") + assert receive_line(sock, buf) == "OK" + for request in [f"WANT {bad_hash}", "WANT ../../config", "REF ../../config"]: + send_line(sock, request) + assert receive_line(sock, buf).startswith("ERR ") + send_line(sock, f"WANT {good_hash}") + assert receive_line(sock, buf) == "OBJ blob 4" + assert recv_exact(sock, buf, 4) == b"good" + send_line(sock, "DONE") diff --git a/tests/test_week3_integration.py b/tests/test_week3_integration.py new file mode 100644 index 0000000..bbc9b69 --- /dev/null +++ b/tests/test_week3_integration.py @@ -0,0 +1,81 @@ +import threading + +from minigit.cli import main +from minigit.commits import CommitManager +from minigit.index import WorkingTree +from minigit.objects import ObjectStore +from minigit.remote import RemoteClient, RemoteServer + + +def test_cli_commit_status_and_history(tmp_path, monkeypatch, capsys): + monkeypatch.chdir(tmp_path) + assert main(["init"]) == 0 + (tmp_path / "src").mkdir() + file = tmp_path / "src" / "hello.txt" + file.write_text("first") + assert main(["add", "src/hello.txt"]) == 0 + capsys.readouterr() + assert main(["status"]) == 0 + assert capsys.readouterr().out == "staged:\n src/hello.txt\n" + assert main(["commit", "-m", "first"]) == 0 + manager = CommitManager(tmp_path) + first = manager.read_ref("main") + first_tree = manager.get_head_tree() + entries = WorkingTree(tmp_path).read_tree_entries(first_tree) + assert [e.path for e in entries] == ["src/hello.txt"] + assert manager.store.read_object(entries[0].hash) == ("blob", b"first") + capsys.readouterr() + assert main(["status"]) == 0 + assert capsys.readouterr().out == "clean\n" + file.write_text("second") + assert main(["add", "src/hello.txt"]) == 0 + file.write_text("third") + assert main(["status"]) == 0 + assert capsys.readouterr().out == "staged:\n src/hello.txt\nnot staged:\n src/hello.txt\n" + assert main(["commit", "-m", "second"]) == 0 + second = manager.read_ref("main") + assert manager.read_commit(second).parents == [first] + assert manager.store.read_object(entries[0].hash) == ("blob", b"first") + capsys.readouterr() + assert main(["log"]) == 0 + assert capsys.readouterr().out == f"{second[:7]} second\n{first[:7]} first\n" + + +def test_fetch_real_merge_history_and_nested_trees(tmp_path): + source = tmp_path / "source" + source.mkdir() + (source / "src").mkdir() + file = source / "src" / "data.bin" + file.write_bytes(b"binary\x00\ncontent") + wt = WorkingTree(source) + wt.stage_file("src/data.bin") + manager = CommitManager(source) + first_tree = wt.build_tree_from_index() + base = manager.create_commit(first_tree, [], "Test ", "base") + left = manager.create_commit(first_tree, [base], "Test ", "left") + file.write_bytes(b"changed") + wt.stage_file("src/data.bin") + second_tree = wt.build_tree_from_index() + right = manager.create_commit(second_tree, [base], "Test ", "right") + tip = manager.create_commit(second_tree, [left, right], "Test ", "merge") + source_client = RemoteClient(source) + reachable = source_client.collect_reachable("main") + assert {base, left, right, tip, first_tree, second_tree} <= reachable + assert len(reachable) == 10 # four commits, four trees, two blobs + server = RemoteServer(source, token="test-token") + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + destination = tmp_path / "destination" + client = RemoteClient(destination) + try: + client.fetch_objects(f"127.0.0.1:{server.port}", sorted(reachable), "test-token") + finally: + server.close() + thread.join(timeout=2) + assert not thread.is_alive() + store = ObjectStore(destination) + for obj_hash in reachable: + assert store.read_object(obj_hash) == manager.store.read_object(obj_hash) + client.commits.write_ref("main", tip) + assert client.collect_reachable("main") == reachable + assert client.commits.log() == manager.log()