Skip to content
Draft
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
66 changes: 27 additions & 39 deletions servers/lean_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,18 +30,20 @@
INSTALL_PATH_ID = hashlib.sha256(os.fsencode(PACKAGE_ROOT)).hexdigest()[:10]


# Code whose changes require replacing a live runtime daemon.
_RUNTIME_SOURCES = (
Path(__file__).resolve(),
PACKAGE_ROOT / "servers" / "lean_runtime.py",
PACKAGE_ROOT / "servers" / "lsp" / "server.py",
PACKAGE_ROOT / "servers" / "repl" / "core.py",
PACKAGE_ROOT / "servers" / "repl" / "pool.py",
)


def _build_id() -> str:
"""Fingerprint code that can change persistent runtime behavior."""
digest = hashlib.sha256()
runtime_files = (
PACKAGE_ROOT / "servers" / "__init__.py",
Path(__file__).resolve(),
PACKAGE_ROOT / "servers" / "lean_runtime.py",
PACKAGE_ROOT / "servers" / "lsp" / "server.py",
PACKAGE_ROOT / "servers" / "repl" / "core.py",
PACKAGE_ROOT / "servers" / "repl" / "pool.py",
)
for path in runtime_files:
for path in (PACKAGE_ROOT / "servers" / "__init__.py", *_RUNTIME_SOURCES):
try:
digest.update(path.read_bytes())
except OSError:
Expand All @@ -51,15 +53,8 @@ def _build_id() -> str:

def _build_generation() -> int:
"""Order in-place builds so an older live wrapper cannot replace a newer one."""
candidates = (
Path(__file__).resolve(),
PACKAGE_ROOT / "servers" / "lean_runtime.py",
PACKAGE_ROOT / "servers" / "lsp" / "server.py",
PACKAGE_ROOT / "servers" / "repl" / "core.py",
PACKAGE_ROOT / "servers" / "repl" / "pool.py",
)
mtimes: list[int] = []
for path in candidates:
for path in _RUNTIME_SOURCES:
try:
mtimes.append(path.stat().st_mtime_ns)
except OSError:
Expand Down Expand Up @@ -463,7 +458,6 @@ def _request_once(
raise LeanRuntimeProtocolError("Lean runtime request exceeds the message limit")

connection = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
dispatched = False
try:
connection.settimeout(self.connect_timeout)
try:
Expand All @@ -473,36 +467,30 @@ def _request_once(
f"Lean runtime is not listening at {self.paths.socket}"
) from error
except OSError as error:
if error.errno in {2, 61, 111}:
raise LeanRuntimeUnavailable(
f"Lean runtime is not listening at {self.paths.socket}"
) from error
raise LeanRuntimeError(f"cannot connect to Lean runtime: {error}") from error

connection.settimeout(response_timeout or self.response_timeout)
# From this point onward, any failure is ambiguous: the daemon may
# have received the request. Never auto-replay Lean execution.
dispatched = True
connection.sendall(payload)
raw = self._read_line(connection)
except socket.timeout as error:
phase = "response" if dispatched else "connection"
raise LeanRuntimeError(f"timed out waiting for Lean runtime {phase}") from error
except LeanRuntimeError:
raise
except OSError as error:
if not dispatched:
raise LeanRuntimeUnavailable(
f"Lean runtime is not listening at {self.paths.socket}"
try:
connection.sendall(payload)
raw = self._read_line(connection)
except socket.timeout as error:
raise LeanRuntimeError("timed out waiting for Lean runtime response") from error
except OSError as error:
raise LeanRuntimeError(
"connection to Lean runtime closed after request dispatch; the request was not retried"
) from error
raise LeanRuntimeError(
"connection to Lean runtime closed after request dispatch; the request was not retried"
except OSError as error:
# Only settimeout() reaches here; nothing has been sent yet.
raise LeanRuntimeUnavailable(
f"Lean runtime is not listening at {self.paths.socket}"
) from error
finally:
connection.close()

try:
response = json.loads(raw)
response = json.loads(raw.decode("utf-8"))
except (UnicodeDecodeError, json.JSONDecodeError) as error:
raise LeanRuntimeProtocolError("Lean runtime returned invalid JSON") from error
if not isinstance(response, dict):
Expand All @@ -523,7 +511,7 @@ def _request_once(
raise LeanRuntimeRemoteError(f"{error_type}: {message}")

@staticmethod
def _read_line(connection: socket.socket) -> str:
def _read_line(connection: socket.socket) -> bytes:
data = bytearray()
while len(data) <= MAX_MESSAGE_BYTES:
chunk = connection.recv(min(65536, MAX_MESSAGE_BYTES + 1 - len(data)))
Expand All @@ -534,5 +522,5 @@ def _read_line(connection: socket.socket) -> str:
if newline >= 0:
if data[newline + 1 :]:
raise LeanRuntimeProtocolError("Lean runtime returned trailing response data")
return bytes(data[:newline]).decode("utf-8")
return bytes(data[:newline])
raise LeanRuntimeProtocolError("Lean runtime response exceeds the message limit")
26 changes: 26 additions & 0 deletions tests/test_shared_lean_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -886,6 +886,32 @@ def close(self):
client.request("repl.run", {"project_dir": "/lean", "code": "#check Nat"})


def test_non_utf8_response_is_a_protocol_error(runtime_dir, monkeypatch):
from servers import lean_client

class NonUtf8Socket:
def settimeout(self, timeout):
pass

def connect(self, path):
pass

def sendall(self, payload):
pass

def recv(self, size):
return b"\xff\n"

def close(self):
pass

client = LeanRuntimeClient(socket_path=runtime_dir / "fake.sock")
monkeypatch.setattr(lean_client.socket, "socket", lambda *args: NonUtf8Socket())

with pytest.raises(lean_client.LeanRuntimeProtocolError, match="invalid JSON"):
client.request("repl.run", {"project_dir": "/lean", "code": "#check Nat"})


@pytest.mark.parametrize(
("name", "value", "match"),
[
Expand Down
Loading