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
57 changes: 44 additions & 13 deletions python/leotower/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,22 @@ def __repr__(self):
return f"Goal({self.hyps!r} ⊢ {self.ty})"


class LeanError(RuntimeError):
"""Base class for leotower Lean-operation errors.

A subclass of :class:`RuntimeError`, so existing ``except RuntimeError``
handlers keep working unchanged. Raised by :class:`Repl` operations
that fail outside of tactic application: :meth:`Repl.set_goal`,
:meth:`Repl.check`, :meth:`Repl.inspect`, and :meth:`Repl.run_cmd`.
"""


class TacticError(LeanError):
"""A tactic failed to parse, elaborate, or apply; the session and the
replay state stay usable.
"""


class Repl:
"""A LeanDojo-style replay session over the embedded Lean runtime.

Expand All @@ -120,20 +136,26 @@ def __init__(self, module: str = "Lean"):
# -- state management ---------------------------------------------------
def set_goal(self, type_str: str) -> int:
"""Set the root goal from a term string; returns state 0."""
return self._repl.set_goal(type_str)
try:
return self._repl.set_goal(type_str)
except RuntimeError as e:
raise LeanError(str(e)) from e

def run_tac(self, state: int, tactic: str, goal_idx: int = 0) -> int:
"""Apply ``tactic`` to the ``goal_idx``-th goal of ``state`` (default
0); returns the new state id. Multi-goal states (from ``induction``,
``split``, ``cases``) keep their unworked goals in the new state, so
a proof can advance goal by goal in any order. Invalid tactics
raise :class:`RuntimeError` (the interpreter and the replay state
raise :class:`TacticError` (the interpreter and the replay state
stay intact)."""
return self._repl.run_tac(state, tactic, goal_idx)
try:
return self._repl.run_tac(state, tactic, goal_idx)
except RuntimeError as e:
raise TacticError(str(e)) from e

def try_run_tac(self, state: int, tactic: str, goal_idx: int = 0) -> "tuple[int, bool]":
"""Non-raising variant of :meth:`run_tac`: returns
``(state_id, success)`` instead of raising :class:`RuntimeError`.
``(state_id, success)`` instead of raising :class:`TacticError`.

On success the new state id and ``True`` are returned. On failure
(unknown state, out-of-range goal, or tactic parse/elaboration/run
Expand All @@ -152,12 +174,15 @@ def run_cmd(self, cmd: str) -> None:
The command is parsed with Lean's real parser and elaborated by the
embedded frontend (``Lean.Elab.Command.elabCommandTopLevel``); the
resulting environment is installed for subsequent calls. Commands
that fail elaboration raise :class:`RuntimeError` (the session stays
that fail elaboration raise :class:`LeanError` (the session stays
usable). Commands do not create replay states and return nothing —
use :meth:`inspect` / :meth:`check` for declaration and term
queries, and only run environment-mutating commands here.
"""
self._repl.run_cmd(cmd)
try:
self._repl.run_cmd(cmd)
except RuntimeError as e:
raise LeanError(str(e)) from e

# -- goal queries -------------------------------------------------------
def get_num_goals(self, state: int) -> int:
Expand All @@ -175,7 +200,7 @@ def run_tacs(self, state: int, tactics: "list[str]", goal_idx: int = 0) -> int:
starting from ``state``; returns the final state id. This is the
replay/RL loop idiom — apply ``[t1, t2, ...]`` without threading
intermediate state ids by hand. Tactics are applied left to right;
the first one to fail raises :class:`RuntimeError` and the returned
the first one to fail raises :class:`TacticError` and the returned
state is that of the last successful tactic (the session stays
usable). If ``tactics`` is empty, ``state`` is returned unchanged.
"""
Expand All @@ -188,7 +213,7 @@ def try_run_tacs(self, state: int, tactics: "list[str]", goal_idx: int = 0) -> "
"""Non-raising variant of :meth:`run_tacs`: apply a sequence of
tactics in order to the ``goal_idx``-th goal, starting from
``state``, and return ``(state_id, success)`` instead of raising
:class:`RuntimeError`.
:class:`TacticError`.

If every tactic succeeds, ``(final_state, True)`` is returned. If
one fails midway, the state after the last successful tactic and
Expand Down Expand Up @@ -236,13 +261,16 @@ def check(self, term: str, state: int | None = None, goal_idx: int = 0) -> str:
Names are resolved at the meta level: use fully qualified names
(command-level scopes such as ``open`` do not apply). Elaboration
failures (unknown identifiers, type errors) raise
:class:`RuntimeError` with Lean's error message; the replay session
:class:`LeanError` with Lean's error message; the replay session
is not modified.

>>> repl.check("Nat.add")
'Nat.add : Nat → Nat → Nat'
"""
return self._repl.check(term, state, goal_idx)
try:
return self._repl.check(term, state, goal_idx)
except RuntimeError as e:
raise LeanError(str(e)) from e

def inspect(self, name: str) -> str:
"""``#print``-style query: show the declaration's kind, type, and
Expand All @@ -252,15 +280,18 @@ def inspect(self, name: str) -> str:
>>> repl.inspect("Nat.add")
'def Nat.add : Nat → Nat → Nat := ...'

Unknown declarations raise :class:`RuntimeError`. Declarations
Unknown declarations raise :class:`LeanError`. Declarations
created with :meth:`run_cmd` (``def``, ``axiom``, ...) are visible
here.
"""
return self._repl.inspect(name)
try:
return self._repl.inspect(name)
except RuntimeError as e:
raise LeanError(str(e)) from e

# -- environment queries ------------------------------------------------
def env_has_const(self, name: str) -> bool:
return self._repl.env_has_const(name)


__all__ += ["Repl", "Goal"]
__all__ += ["Repl", "Goal", "LeanError", "TacticError"]
69 changes: 68 additions & 1 deletion tests/test_repl.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@

import pytest

from leotower import Repl
from leotower import Repl, LeanError, TacticError

# Windows: constructing Repl() aborts the whole process (misaligned pointer
# dereference in leo3-ffi, exit 127, no Python traceback) — tracked in W-395.
Expand Down Expand Up @@ -729,3 +729,70 @@ def test_num_states_counts():
assert repl.num_states() == 1
s1 = repl.run_tac(s0, "intro n")
assert repl.num_states() == 2


# ============================================================================
# Exception hierarchy: LeanError / TacticError
# ============================================================================


def test_exception_hierarchy():
"""TacticError refines LeanError; both refine RuntimeError so existing
``except RuntimeError`` handlers keep working (backward compatibility)."""
assert issubclass(LeanError, RuntimeError)
assert issubclass(TacticError, LeanError)


def test_tactic_failures_raise_tactic_error():
"""Every tactic-failure path raises TacticError with the original
message and cause preserved."""
repl = Repl()
s0 = repl.set_goal(ADD_COMM)
# Type error.
with pytest.raises(TacticError, match="tactic error") as excinfo:
repl.run_tac(s0, "exact 42")
# The original message and the underlying cause are preserved.
assert isinstance(excinfo.value.__cause__, RuntimeError)
# Parse error — a bare ``except RuntimeError`` still catches it
# (backward compatibility) and the concrete type is TacticError.
with pytest.raises(RuntimeError, match="tactic parse error") as excinfo:
repl.run_tac(s0, "this is not a tactic !!!")
assert isinstance(excinfo.value, TacticError)
# Unknown state.
with pytest.raises(TacticError, match="unknown state"):
repl.run_tac(99, "intro n")
# run_tacs propagates the TacticError of the failing element.
s1 = repl.run_tac(s0, "intro n m")
with pytest.raises(TacticError):
# "intro n m" already introduced both vars; re-introducing fails.
repl.run_tacs(s1, ["intro n m", "rfl"])


def test_non_tactic_failures_raise_lean_error():
"""set_goal / run_cmd / check / inspect failures raise LeanError."""
repl = Repl()
with pytest.raises(LeanError):
repl.set_goal("this is not a term @@@")
with pytest.raises(LeanError, match="command parse error"):
repl.run_cmd("this is not a command @@@")
with pytest.raises(LeanError, match="Unknown identifier"):
repl.check("no_such_decl_xyz")
with pytest.raises(LeanError, match="unknown constant"):
repl.inspect("no_such_decl_xyz")


def test_every_error_is_caught_by_except_runtime_error():
"""Backward compatibility across the whole hierarchy: each wrapped
failure is still caught by a plain ``except RuntimeError`` handler."""
repl = Repl()
s0 = repl.set_goal(ADD_COMM)
for call in (
lambda: repl.run_tac(s0, "this is not a tactic !!!"),
lambda: repl.run_tacs(s0, ["this is not a tactic !!!"]),
lambda: repl.set_goal("this is not a term @@@"),
lambda: repl.run_cmd("this is not a command @@"),
lambda: repl.check("no_such_decl_xyz"),
lambda: repl.inspect("no_such_decl_xyz"),
):
with pytest.raises(RuntimeError):
call()
Loading