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
158 changes: 152 additions & 6 deletions minigit/remote.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,10 +12,37 @@
# from .objects import ObjectStore (Module 1 & 3)
# from .commits import CommitManager
import os
import socket

from minigit.errors import NetworkProtocolError


def send_line(sock, text: str) -> None:
"""Send one line of the wire protocol,"""

sock.sendall((text + "\n").encode())


def receive_line(sock, buf: bytearray) -> str:
"""Read one `\\n`-terminated line from `sock`.

`buf` is the connection's leftover-bytes buffer, owned by the caller and
reused across calls on the same socket: a `recv()` can return two lines
at once (the second one waits here for the next call) or half a line
(we keep reading until the rest arrives).
"""

while b"\n" not in buf:
chunk = sock.recv(4096)
if not chunk:
raise NetworkProtocolError("connection closed mid-line")
buf.extend(chunk)

line, _, rest = buf.partition(b"\n")
buf[:] = rest
return line.decode()


class RemoteClient:
"""Push and pull commits between two minigit repos over a TCP connection."""

Expand Down Expand Up @@ -55,10 +82,27 @@ def push(self, remote_address: str, branch: str, token: str) -> None:
host, port = self._parse_address(remote_address)
if len(token) == 0:
raise NetworkProtocolError("push needs a token: pass --token")
print(f"push: would push {branch} to {host}:{port}")

# connect over TCP
# ask remote for its current hash for <branch>
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}")

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")
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
Expand All @@ -70,8 +114,96 @@ def pull(self, remote_address: str, branch: str, token: str) -> None:
host, port = self._parse_address(remote_address)
if len(token) == 0:
raise NetworkProtocolError("pull needs a token: pass --token")
print(f"pull: would pull {branch} from {host}:{port}")
# Week 6 - same exchange in reverse

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}")

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 - same exchange in reverse")
finally:
sock.close()


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):
self.repo_path = repo_path
self.token = token
self.host = host

self._sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
self._sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
self._sock.bind((host, port))
self._sock.listen()
# accept() polls this instead of blocking forever: closing the socket
# from another thread isn't guaranteed to unblock a pending accept()
# on every platform (it does on macOS, not reliably on Linux).
self._sock.settimeout(0.5)
self.port = self._sock.getsockname()[1]
self._closed = False

def close(self) -> None:
"""Stop accepting connections. `serve_forever()` notices within one poll interval."""

self._closed = True
self._sock.close()

def serve_forever(self) -> None:
"""Accept connections and handle them one at a time until `close()` is called."""

while not self._closed:
try:
conn, _ = self._sock.accept()
except TimeoutError:
continue # no connection yet - check self._closed and try again
except OSError:
return # listening socket was closed - shut down

try:
self._handle_client(conn)
except (NetworkProtocolError, OSError):
pass # a bad client must not take down the server
finally:
conn.close()

def _handle_client(self, conn) -> None:
"""Run one client's AUTH + REF handshake."""

buf = bytearray()

line = receive_line(conn, buf)
command, _, value = line.partition(" ")
if command != "AUTH" or value != self.token:
send_line(conn, "ERR bad auth")
return
send_line(conn, "OK")

line = receive_line(conn, buf)
command, _, branch = line.partition(" ")
if command != "REF":
send_line(conn, "ERR expected REF")
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 = "-"
send_line(conn, f"REF {branch} {commit_hash}")


# Wire protocol (draft only - Week 2 makes this real):
Expand All @@ -86,7 +218,7 @@ def pull(self, remote_address: str, branch: str, token: str) -> None:


def register_subcommands(subparsers) -> None:
"""Register the `push` and `pull` subcommands with the CLI parser."""
"""Register the `push`, `pull`, and `serve` subcommands with the CLI parser."""

push_parser = subparsers.add_parser("push", help="push a branch to a remote")
push_parser.add_argument("address")
Expand All @@ -100,6 +232,11 @@ def register_subcommands(subparsers) -> None:
pull_parser.add_argument("--token", default="")
pull_parser.set_defaults(handler=cmd_pull)

serve_parser = subparsers.add_parser("serve", help="serve this repo to push/pull clients")
serve_parser.add_argument("--port", type=int, required=True)
serve_parser.add_argument("--token", default="")
serve_parser.set_defaults(handler=cmd_serve)


def cmd_push(args) -> int:
"""Handle `minigit push` from the CLI."""
Expand All @@ -113,3 +250,12 @@ def cmd_pull(args) -> int:

RemoteClient().pull(args.address, args.branch, args.token)
return 0


def cmd_serve(args) -> int:
"""Handle `minigit serve` from the CLI."""

server = RemoteServer(port=args.port, token=args.token)
print(f"listening on {server.host}:{server.port}")
server.serve_forever()
return 0
72 changes: 69 additions & 3 deletions tests/test_remote.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,11 @@
import threading

import pytest

from minigit.errors import NetworkProtocolError
from minigit.remote import RemoteClient
from minigit.remote import RemoteClient, RemoteServer, receive_line

KNOWN_HASH = "a" * 40


class FakeObjectStore:
Expand All @@ -16,6 +20,23 @@ def make_client():
return RemoteClient(store=FakeObjectStore(), commits=FakeCommitManager())


@pytest.fixture
def remote_server(tmp_path):
refs_dir = tmp_path / ".minigit" / "refs" / "heads"
refs_dir.mkdir(parents=True)
(refs_dir / "main").write_text(KNOWN_HASH + "\n")

server = RemoteServer(repo_path=str(tmp_path), token="tok", port=0)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
server._thread = thread # test-only: lets a test wait for a manual close() to land

yield server

server.close()
thread.join(timeout=2)


def test_parse_address_valid():
client = make_client()
assert client._parse_address("127.0.0.1:9418") == ("127.0.0.1", 9418)
Expand Down Expand Up @@ -51,6 +72,51 @@ def test_push_empty_token_raises():
client.push("127.0.0.1:9418", "main", "")


def test_push_valid_address_and_token_does_not_raise():
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")


def test_push_wrong_token_raises(remote_server):
client = make_client()
with pytest.raises(NetworkProtocolError):
client.push(f"127.0.0.1:{remote_server.port}", "main", "wrong")


def test_push_existing_branch_prints_matching_hash(remote_server, capsys):
client = make_client()
client.push("127.0.0.1:9418", "main", "sometoken")
client.push(f"127.0.0.1:{remote_server.port}", "main", "tok")
assert KNOWN_HASH in capsys.readouterr().out


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(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()
with pytest.raises(NetworkProtocolError):
client.push(f"127.0.0.1:{port}", "main", "tok")


class FakeSocket:
"""Hands back pre-scripted chunks instead of reading a real socket."""

def __init__(self, chunks):
self._chunks = list(chunks)

def recv(self, size):
return self._chunks.pop(0) if self._chunks else b""


def test_receive_line_splits_two_lines_from_one_packet():
sock = FakeSocket([b"AUTH tok\nREF main\n"])
buf = bytearray()

assert receive_line(sock, buf) == "AUTH tok"
assert receive_line(sock, buf) == "REF main"
Loading