File size: 14,716 Bytes
32c0c6c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
"""Batched Lean verifier (Phase 3).

One `lean --json` invocation verifies a whole batch of candidate proof attempts:
each candidate is emitted as its own theorem in a single generated file, and the
per-theorem diagnostics (closed / invalid / resulting goals) are parsed back.

Measured on this machine: ~5 s per invocation for Mathlib import, and the
per-candidate cost is essentially free (24 candidates in 5.2 s vs 4 in 5.8 s),
which is what makes search affordable without a REPL.
"""
import json
import os
import re
import subprocess
import tempfile
import threading
import time

HOME = os.path.expanduser('~')
# Overridable so this package can run outside the author's machine layout:
#   export LEANOAR_LEAN_BIN=/path/to/lean        (the `lean` binary of the 4.35.0-rc2 toolchain)
#   export LEANOAR_MATHLIB=/path/to/mathlib4     (checkout with a prebuilt .lake cache)
LEAN_BIN = os.environ.get('LEANOAR_LEAN_BIN',
                          f'{HOME}/leanprover/toolchains/lean-4.35.0-rc2-linux/bin/lean')
MATHLIB = os.environ.get('LEANOAR_MATHLIB', f'{HOME}/leanprover/mathlib4')


def lean_env():
    env = dict(os.environ)
    pkgs = []
    pdir = f'{MATHLIB}/.lake/packages'
    if os.path.isdir(pdir):
        for p in sorted(os.listdir(pdir)):
            lib = f'{pdir}/{p}/.lake/build/lib/lean'
            if os.path.isdir(lib):
                pkgs.append(lib)
    env['LEAN_PATH'] = ':'.join([f'{MATHLIB}/.lake/build/lib/lean'] + pkgs)
    env['PATH'] = os.path.dirname(LEAN_BIN) + ':' + env.get('PATH', '')
    return env


def split_decl(statement):
    """`theorem foo (x : Nat) : x = x := sorry` -> ('foo', ' (x : Nat) : x = x')

    The trailing `:= sorry` / `:= by sorry` is dropped so we can append `:= by ...`.
    """
    s = statement.strip()
    s = re.sub(r':=\s*(by\s+)?sorry\s*$', '', s).rstrip()
    m = re.match(r'(?:theorem|lemma|example)\s+([A-Za-z0-9_.\'«»!?]+)', s)
    if m:
        name = m.group(1)
        rest = s[m.end():]
    else:
        name = None
        rest = s
    return name, rest


def indent(tactic, first=2, cont=4):
    lines = [l.strip() for l in tactic.strip().split('\n') if l.strip()]
    if not lines:
        return ''
    out = [' ' * first + lines[0]]
    for l in lines[1:]:
        out.append(' ' * cont + l)
    return '\n'.join(out)


# Errors that may mean "this file lacked an import" rather than "this tactic is wrong".
# The narrow-import path re-checks such candidates against the full `import Mathlib`.
MISSING_IMPORT_RE = re.compile(
    r'unknown (?:identifier|constant|namespace|declaration|tactic|module|package)'
    r'|object file .* does not exist'
    r'|failed to synthesize', re.I)


def min_imports(verifier, header, statement, tag='minimp'):
    """Modules that suffice to elaborate `statement` on this machine, via Mathlib's `#min_imports`.

    One full `import Mathlib` call per problem (~3.4 s) that pays for itself at once: every
    later verification call then imports ~1.7 s of modules instead of ~3.4 s (measured 3.41 s ->
    1.74 s per call, 1.9x). Returns None when the answer is unusable, in which case callers must
    keep the original header.
    """
    name, rest = split_decl(statement)
    hdr = header.strip()
    if not hdr.startswith('import '):      # miniF2F headers carry only `open ...`
        hdr = 'import Mathlib\n' + hdr
    src = [hdr, '', f'theorem {name or "goal"}_{tag}{rest} := by sorry', '',
           '#min_imports', '']
    try:
        r = verifier.run_file('\n'.join(src) + '\n', name=f'{name or "goal"}_{tag}')
    except Exception:                      # a broken probe must never break the search
        return None
    mods = []
    for m in r['messages']:
        mods += re.findall(r'public import ([A-Za-z0-9_.«»]+)', str(m.get('data', '')))
    mods = list(dict.fromkeys(mods))
    if r['returncode'] != 0 or not any(m.get('severity') == 'information' for m in r['messages']):
        return None                        # the probe itself did not run to completion
    if len(mods) > 12 or 'Mathlib' in mods:
        return None                        # no saving over `import Mathlib`: keep the original
    return mods                            # may be [] (statement needs almost nothing)


def narrowed_header(verifier, header, statement):
    """Narrowed replacement for `import Mathlib`, or None when none can be built.

    The statement's minimal imports (`#min_imports`), always plus `Mathlib.Tactic` because the
    injected automation table (`ring`, `omega`, `simp_all`, ...) and most model tactics need it,
    plus the original `open ...` lines. Measured 3.41 s -> 1.74 s per `lean --json` call.
    """
    mods = min_imports(verifier, header, statement)
    if mods is None:
        return None
    opens = [l for l in header.splitlines() if not l.strip().startswith('import ')]
    imports = [f'import {m}' for m in mods]
    if 'Mathlib.Tactic' not in mods:
        imports.append('import Mathlib.Tactic')
    return '\n'.join(imports + opens).strip()


class LeanBatchVerifier:
    """Runs one `lean --json` per batch of candidate proof attempts."""

    def __init__(self, workdir=None, timeout=600, keep_files=True, verbose=False):
        self.workdir = workdir or f'{HOME}/leanprover/work/search'
        os.makedirs(self.workdir, exist_ok=True)
        self.timeout = timeout
        self.keep_files = keep_files
        self.verbose = verbose
        self._lock = threading.Lock()
        self.n_calls = 0
        self.lean_seconds = 0.0
        self.self_check()

    def self_check(self):
        """Warm the disk cache once and make sure a known proof compiles."""
        t0 = time.time()
        r = self.run_file('import Mathlib\ntheorem _selfcheck : 1 + 1 = 2 := by norm_num\n',
                          name='selfcheck')
        self.warm_seconds = time.time() - t0
        if r['returncode'] != 0 or r['messages']:
            raise RuntimeError(f'verifier self-check failed: {r["messages"][:2]}')

    # ---------- raw ----------
    def run_file(self, source, name='batch', timeout=None):
        path = os.path.join(self.workdir, f'{name}.lean')
        with open(path, 'w') as f:
            f.write(source)
        t0 = time.time()
        try:
            p = subprocess.run([LEAN_BIN, '--json', path],
                               cwd=self.workdir, env=lean_env(), timeout=timeout or self.timeout,
                               capture_output=True, text=True)
            out, rc, timed_out = p.stdout, p.returncode, False
        except subprocess.TimeoutExpired as e:
            out, rc, timed_out = (e.stdout or b'').decode('utf-8', 'replace') if isinstance(e.stdout, bytes) else (e.stdout or ''), -1, True
        dt = time.time() - t0
        with self._lock:
            self.n_calls += 1
            self.lean_seconds += dt
        msgs = []
        for line in out.splitlines():
            line = line.strip()
            if line.startswith('{'):
                try:
                    msgs.append(json.loads(line))
                except json.JSONDecodeError:
                    pass
        if not self.keep_files:
            os.remove(path)
        return {'messages': msgs, 'returncode': rc, 'timed_out': timed_out,
                'seconds': dt, 'path': path, 'stdout': out if rc == -1 or not msgs else ''}

    # ---------- batch API ----------
    def verify(self, header, items, name='batch', timeout=None):
        """items: list of dicts {name, decl_rest, tactics: [str, ...]}
        (`decl_rest` = the piece after the theorem name, e.g. " (x : Nat) : x = x")

        Returns list of per-item dicts: {status: closed|open|invalid|timeout,
        goals, error, seconds} aligned with `items`.
        """
        src = [header.strip(), '']
        if not src[0].startswith('import '):      # miniF2F headers carry only `open ...`
            src.insert(0, 'import Mathlib')
        nxt = sum(c.count('\n') + 1 for c in src) + 1     # line number of the next chunk
        spans = []
        for i, it in enumerate(items):
            tname = it.get('name') or f'goal_{i}'
            # unique declaration name per item
            uniq = f'{re.sub(r"[^A-Za-z0-9_]", "_", tname)}_v{i}'
            head = f'theorem {uniq}{it["decl_rest"]} := by'
            start = nxt
            src.append(head)
            nxt += head.count('\n') + 1
            for t in it['tactics']:
                block = indent(t)
                src.append(block)
                nxt += block.count('\n') + 1
            end = nxt - 1
            src.append('')
            nxt += 1
            spans.append((start, end, i, uniq))
        source = '\n'.join(src) + '\n'
        r = self.run_file(source, name=name, timeout=timeout)
        out = [{'status': 'timeout', 'goals': None, 'error': 'lean timeout'} for _ in items]
        if r['timed_out']:
            return out, r
        # Assign every message to the LAST item that started at or before its line.
        # (Requiring strict membership mis-assigns syntax errors, which Lean reports
        # on the line after the offending tactic, and silently turns them into
        # "no errors => closed" false positives.)
        starts = [s for s, _e, _i, _u in spans]
        buckets = [[] for _ in items]
        other = []
        for m in r['messages']:
            ln = m.get('pos', {}).get('line', -1)
            idx = None
            for j, s in enumerate(starts):
                if ln >= s:
                    idx = spans[j][2]
                else:
                    break
            if idx is None:
                other.append(m)
            else:
                buckets[idx].append(m)
        for i, msgs in enumerate(buckets):
            errs = [m for m in msgs if m.get('severity') == 'error']
            sorry = [m for m in msgs if 'sorry' in str(m.get('data', '')).lower()]
            hard = [m for m in errs if m.get('kind') != 'Tactic.unsolvedGoals']
            goal_msgs = [m for m in errs if m.get('kind') == 'Tactic.unsolvedGoals']
            if hard:
                out[i] = {'status': 'invalid', 'goals': None,
                          'error': str(hard[0].get('data'))[:400]}
            elif sorry:
                out[i] = {'status': 'invalid', 'goals': None, 'error': 'uses sorry'}
            elif goal_msgs:
                g = str(goal_msgs[-1].get('data', '')).replace('unsolved goals\n', '').strip()
                out[i] = {'status': 'open', 'goals': g, 'error': None}
            else:
                out[i] = {'status': 'maybe-closed', 'goals': None, 'error': None}
        return out, r

    def verify_many(self, header, items, chunk=24, parallel=4, name='batch', timeout=None):
        """Verify `items` in chunks of `chunk`, running up to `parallel` chunks at once.

        Returns (results, raws): results aligned with `items`; raws per chunk. Measured
        motivation: >90% of search wall-clock is Lean elaboration, and every chunk is an
        independent `lean --json` process (~150 MB RSS), so concurrency on a 16-core box
        is the cheapest remaining speed-up.
        """
        chunks = [items[i:i + chunk] for i in range(0, len(items), chunk)]
        if parallel <= 1 or len(chunks) <= 1:
            out, raws = [], []
            for i, c in enumerate(chunks):
                r, raw = self.verify(header, c, name=f'{name}_c{i}', timeout=timeout)
                out.extend(r)
                raws.append(raw)
            return out, raws
        from concurrent.futures import ThreadPoolExecutor, as_completed
        results, raws = [None] * len(chunks), [None] * len(chunks)
        with ThreadPoolExecutor(max_workers=min(parallel, len(chunks))) as ex:
            futs = {ex.submit(self.verify, header, c, f'{name}_c{i}', timeout): i
                    for i, c in enumerate(chunks)}
            for f in as_completed(futs):
                i = futs[f]
                results[i], raws[i] = f.result()
        out = [x for r in results for x in r]
        return out, raws

    def confirm(self, header, name, decl_rest, tactics, timeout=None):
        """Standalone re-check of ONE claimed proof.

        A batch entry is only 'maybe-closed' (absence of errors is not proof of
        success: syntax errors can land outside the item's line span). This compiles
        the proof alone and demands: exit code 0, no error messages, no `sorry`.
        """
        hdr = header.strip() if header.strip().startswith('import ') else f'import Mathlib\n{header.strip()}'
        uniq = f'{re.sub(r"[^A-Za-z0-9_]", "_", name or "goal")}_confirm'
        src = [hdr, '', f'theorem {uniq}{decl_rest} := by']
        for t in tactics:
            src.append(indent(t))
        src.append('')
        r = self.run_file('\n'.join(src) + '\n', name=f'{uniq}_confirm', timeout=timeout)
        msgs = r['messages']
        errs = [m for m in msgs if m.get('severity') == 'error']
        sorry = [m for m in msgs if 'sorry' in str(m.get('data', '')).lower()]
        ok = (r['returncode'] == 0) and not errs and not sorry and not r['timed_out']
        return {'ok': ok, 'returncode': r['returncode'], 'n_errors': len(errs),
                'sorry': bool(sorry), 'seconds': r['seconds'],
                'messages': [str(m.get('data'))[:200] for m in msgs[:3]]}

    def initial_state(self, header, name, decl_rest, tactic='skip'):
        """Get the initial goal text of a theorem by applying a no-op tactic."""
        items = [{'name': name, 'decl_rest': decl_rest, 'tactics': [tactic]}]
        res, raw = self.verify(header, items, name=f'{name}_init')
        r = res[0]
        if r['status'] == 'open':
            return r['goals'], raw
        if r['status'] == 'maybe-closed':
            return None, raw          # theorem closed by a no-op?! (trivially true goal)
        return None, raw


if __name__ == '__main__':
    v = LeanBatchVerifier()
    print(f'warm self-check in {v.warm_seconds:.1f}s')
    header = 'import Mathlib'
    stmt = 'theorem t (n : ℕ) : n + 0 = n := sorry'
    nm, rest = split_decl(stmt)
    goal, raw = v.initial_state(header, nm, rest)
    print('initial goal:', repr(goal))
    items = [{'name': nm, 'decl_rest': rest, 'tactics': [t]}
             for t in ['simp', 'omega', 'bogus_tac', 'norm_num', 'rfl', 'induction n with | zero => simp | succ n ih => simp [Nat.add_succ]']]
    res, raw2 = v.verify(header, items, name='t_batch')
    for it, r in zip(items, res):
        print(f'  {it["tactics"][0][:44]:46s} -> {r["status"]:8s} {(r["error"] or r["goals"] or "")[:60]!r}')
    print(f'lean calls={v.n_calls} lean_seconds={v.lean_seconds:.1f} '
          f'(init {raw["seconds"]:.1f}s + batch {raw2["seconds"]:.1f}s for {len(items)} candidates)')