diff --git a/kmir/src/kmir/_prove.py b/kmir/src/kmir/_prove.py new file mode 100644 index 000000000..bebc3bd4d --- /dev/null +++ b/kmir/src/kmir/_prove.py @@ -0,0 +1,225 @@ +from __future__ import annotations + +import logging +import tempfile +from pathlib import Path +from typing import TYPE_CHECKING + +from pyk.cterm import CTerm +from pyk.kast.inner import KSequence, KVariable, Subst +from pyk.kast.manip import abstract_term_safely, split_config_from +from pyk.kcfg import KCFG +from pyk.proof.reachability import APRProof, APRProver + +from .cargo import cargo_get_smir_json +from .kast import SymbolicMode, make_call_config +from .kmir import KMIR +from .smir import SMIRInfo + +if TYPE_CHECKING: + from typing import Final + + from pyk.kast.inner import KInner + + from .options import ProveRSOpts + + +_LOGGER: Final = logging.getLogger(__name__) + + +def prove_rs(opts: ProveRSOpts) -> APRProof: + if not opts.rs_file.is_file(): + raise ValueError(f'Input file does not exist: {opts.rs_file}') + + label = f'{opts.rs_file.stem}.{opts.start_symbol}' + + if opts.proof_dir is not None: + target_path = opts.proof_dir / label + return _prove_rs(opts, target_path, label) + + with tempfile.TemporaryDirectory() as tmp_dir: + target_path = Path(tmp_dir) + return _prove_rs(opts, target_path, label) + + +def _prove_rs(opts: ProveRSOpts, target_path: Path, label: str) -> APRProof: + if not opts.reload and opts.proof_dir is not None and APRProof.proof_data_exists(label, opts.proof_dir): + _LOGGER.info(f'Reading proof from disc: {opts.proof_dir}, {label}') + apr_proof = APRProof.read_proof_data(opts.proof_dir, label) + + smir_info = SMIRInfo.from_file(target_path / 'smir.json') + kmir = KMIR.from_kompiled_kore( + smir_info, + target_dir=target_path, + extra_module=opts.add_module, + bug_report=opts.bug_report, + symbolic=True, + haskell_target=opts.haskell_target, + llvm_lib_target=opts.llvm_lib_target, + ) + else: + _LOGGER.info(f'Constructing initial proof: {label}') + if opts.smir: + smir_info = SMIRInfo.from_file(opts.rs_file) + else: + smir_info = SMIRInfo(cargo_get_smir_json(opts.rs_file, save_smir=opts.save_smir)) + + smir_info = smir_info.reduce_to(opts.start_symbol) + # Report whether the reduced call graph includes any functions without MIR bodies + missing_body_syms = [ + sym + for sym, item in smir_info.items.items() + if 'MonoItemFn' in item['mono_item_kind'] and item['mono_item_kind']['MonoItemFn'].get('body') is None + ] + has_missing = len(missing_body_syms) > 0 + _LOGGER.info(f'Reduced items table size {len(smir_info.items)}') + if has_missing: + _LOGGER.info(f'missing-bodies-present={has_missing} count={len(missing_body_syms)}') + _LOGGER.debug(f'Missing-body function symbols (first 5): {missing_body_syms[:5]}') + + kmir = KMIR.from_kompiled_kore( + smir_info, + target_dir=target_path, + extra_module=opts.add_module, + bug_report=opts.bug_report, + symbolic=True, + haskell_target=opts.haskell_target, + llvm_lib_target=opts.llvm_lib_target, + ) + + apr_proof = apr_proof_from_smir( + kmir, + label, + smir_info, + start_symbol=opts.start_symbol, + proof_dir=opts.proof_dir, + ) + if apr_proof.proof_dir is not None and (apr_proof.proof_dir / label).is_dir(): + smir_info.dump(apr_proof.proof_dir / apr_proof.id / 'smir.json') + + if apr_proof.passed: + return apr_proof + + cut_point_rules = _cut_point_rules( + break_on_calls=opts.break_on_calls, + break_on_function_calls=opts.break_on_function_calls, + break_on_intrinsic_calls=opts.break_on_intrinsic_calls, + break_on_thunk=opts.break_on_thunk or opts.terminate_on_thunk, # must break for terminal rule + break_every_statement=opts.break_every_statement, + break_on_terminator_goto=opts.break_on_terminator_goto, + break_on_terminator_switch_int=opts.break_on_terminator_switch_int, + break_on_terminator_return=opts.break_on_terminator_return, + break_on_terminator_call=opts.break_on_terminator_call, + break_on_terminator_assert=opts.break_on_terminator_assert, + break_on_terminator_drop=opts.break_on_terminator_drop, + break_on_terminator_unreachable=opts.break_on_terminator_unreachable, + break_every_terminator=opts.break_every_terminator, + break_every_step=opts.break_every_step, + ) + + with kmir.kcfg_explore(label, terminate_on_thunk=opts.terminate_on_thunk) as kcfg_explore: + prover = APRProver( + kcfg_explore, + execute_depth=opts.max_depth, + cut_point_rules=cut_point_rules, + ) + prover.advance_proof( + apr_proof, + max_iterations=opts.max_iterations, + fail_fast=opts.fail_fast, + maintenance_rate=opts.maintenance_rate, + ) + return apr_proof + + +def apr_proof_from_smir( + kmir: KMIR, + id: str, + smir_info: SMIRInfo, + *, + start_symbol: str = 'main', + proof_dir: Path | None = None, +) -> APRProof: + lhs_config, constraints = make_call_config( + kmir.definition, + smir_info=smir_info, + start_symbol=start_symbol, + mode=SymbolicMode(), + ) + lhs = CTerm(lhs_config, constraints) + + var_config, var_subst = split_config_from(lhs_config) + _rhs_subst: dict[str, KInner] = { + v_name: abstract_term_safely(KVariable('_'), base_name=v_name) for v_name in var_subst + } + _rhs_subst['K_CELL'] = KSequence([KMIR.Symbols.END_PROGRAM]) + rhs = CTerm(Subst(_rhs_subst)(var_config)) + kcfg = KCFG() + init_node = kcfg.create_node(lhs) + target_node = kcfg.create_node(rhs) + return APRProof(id, kcfg, [], init_node.id, target_node.id, {}, proof_dir=proof_dir) + + +def _cut_point_rules( + *, + break_on_calls: bool, + break_on_function_calls: bool, + break_on_intrinsic_calls: bool, + break_on_thunk: bool, + break_every_statement: bool, + break_on_terminator_goto: bool, + break_on_terminator_switch_int: bool, + break_on_terminator_return: bool, + break_on_terminator_call: bool, + break_on_terminator_assert: bool, + break_on_terminator_drop: bool, + break_on_terminator_unreachable: bool, + break_every_terminator: bool, + break_every_step: bool, +) -> list[str]: + cut_point_rules = [] + if break_on_thunk: + cut_point_rules.append('RT-DATA.thunk') + if break_every_statement or break_every_step: + cut_point_rules.extend( + [ + 'KMIR-CONTROL-FLOW.execStmt', + 'KMIR-CONTROL-FLOW.execStmt.union', + ] + ) + if break_on_terminator_goto or break_every_terminator or break_every_step: + cut_point_rules.append('KMIR-CONTROL-FLOW.termGoto') + if break_on_terminator_switch_int or break_every_terminator or break_every_step: + cut_point_rules.append('KMIR-CONTROL-FLOW.termSwitchInt') + if break_on_terminator_return or break_every_terminator or break_every_step: + cut_point_rules.extend( + [ + 'KMIR-CONTROL-FLOW.termReturnSome', + 'KMIR-CONTROL-FLOW.termReturnNone', + 'KMIR-CONTROL-FLOW.endprogram-return', + 'KMIR-CONTROL-FLOW.endprogram-no-return', + ] + ) + if ( + break_on_intrinsic_calls + or break_on_calls + or break_on_terminator_call + or break_every_terminator + or break_every_step + ): + cut_point_rules.append('KMIR-CONTROL-FLOW.termCallIntrinsic') + if ( + break_on_function_calls + or break_on_calls + or break_on_terminator_call + or break_every_terminator + or break_every_step + ): + cut_point_rules.append('KMIR-CONTROL-FLOW.termCallFunction') + if break_on_terminator_assert or break_every_terminator or break_every_step: + cut_point_rules.append('KMIR-CONTROL-FLOW.termAssert') + if break_on_terminator_drop or break_every_terminator or break_every_step: + cut_point_rules.append('KMIR-CONTROL-FLOW.termDrop') + if break_on_terminator_unreachable or break_every_terminator or break_every_step: + cut_point_rules.append('KMIR-CONTROL-FLOW.termUnreachable') + return cut_point_rules diff --git a/kmir/src/kmir/kmir.py b/kmir/src/kmir/kmir.py index 84cd8eaf6..ab913ec31 100644 --- a/kmir/src/kmir/kmir.py +++ b/kmir/src/kmir/kmir.py @@ -1,42 +1,41 @@ from __future__ import annotations import logging -import tempfile from contextlib import contextmanager from functools import cached_property -from pathlib import Path from typing import TYPE_CHECKING from pyk.cli.utils import bug_report_arg -from pyk.cterm import CTerm, cterm_symbolic -from pyk.kast.inner import KApply, KLabel, KSequence, KSort, KToken, KVariable, Subst -from pyk.kast.manip import abstract_term_safely, split_config_from -from pyk.kcfg import KCFG +from pyk.cterm import cterm_symbolic +from pyk.kast.inner import KApply, KLabel, KSequence, KSort, KToken from pyk.kcfg.explore import KCFGExplore from pyk.kcfg.semantics import DefaultSemantics from pyk.kcfg.show import NodePrinter from pyk.ktool.kprove import KProve from pyk.ktool.krun import KRun -from pyk.proof.reachability import APRProof, APRProver from pyk.proof.show import APRProofNodePrinter -from .cargo import cargo_get_smir_json -from .kast import ConcreteMode, RandomMode, SymbolicMode, make_call_config +from .kast import ConcreteMode, RandomMode, make_call_config from .kparse import KParse from .parse.parser import Parser from .smir import SMIRInfo if TYPE_CHECKING: from collections.abc import Iterator + from pathlib import Path from typing import Final + from pyk.cterm import CTerm from pyk.cterm.show import CTermShow from pyk.kast.inner import KInner + from pyk.kcfg import KCFG from pyk.kore.syntax import Pattern + from pyk.proof.reachability import APRProof from pyk.utils import BugReport from .options import DisplayOpts, ProveRSOpts + _LOGGER: Final = logging.getLogger(__name__) @@ -53,70 +52,6 @@ def __init__( KParse.__init__(self, definition_dir) self.llvm_library_dir = llvm_library_dir - @staticmethod - def cut_point_rules( - break_on_calls: bool, - break_on_function_calls: bool, - break_on_intrinsic_calls: bool, - break_on_thunk: bool, - break_every_statement: bool, - break_on_terminator_goto: bool, - break_on_terminator_switch_int: bool, - break_on_terminator_return: bool, - break_on_terminator_call: bool, - break_on_terminator_assert: bool, - break_on_terminator_drop: bool, - break_on_terminator_unreachable: bool, - break_every_terminator: bool, - break_every_step: bool, - ) -> list[str]: - cut_point_rules = [] - if break_on_thunk: - cut_point_rules.append('RT-DATA.thunk') - if break_every_statement or break_every_step: - cut_point_rules.extend( - [ - 'KMIR-CONTROL-FLOW.execStmt', - 'KMIR-CONTROL-FLOW.execStmt.union', - ] - ) - if break_on_terminator_goto or break_every_terminator or break_every_step: - cut_point_rules.append('KMIR-CONTROL-FLOW.termGoto') - if break_on_terminator_switch_int or break_every_terminator or break_every_step: - cut_point_rules.append('KMIR-CONTROL-FLOW.termSwitchInt') - if break_on_terminator_return or break_every_terminator or break_every_step: - cut_point_rules.extend( - [ - 'KMIR-CONTROL-FLOW.termReturnSome', - 'KMIR-CONTROL-FLOW.termReturnNone', - 'KMIR-CONTROL-FLOW.endprogram-return', - 'KMIR-CONTROL-FLOW.endprogram-no-return', - ] - ) - if ( - break_on_intrinsic_calls - or break_on_calls - or break_on_terminator_call - or break_every_terminator - or break_every_step - ): - cut_point_rules.append('KMIR-CONTROL-FLOW.termCallIntrinsic') - if ( - break_on_function_calls - or break_on_calls - or break_on_terminator_call - or break_every_terminator - or break_every_step - ): - cut_point_rules.append('KMIR-CONTROL-FLOW.termCallFunction') - if break_on_terminator_assert or break_every_terminator or break_every_step: - cut_point_rules.append('KMIR-CONTROL-FLOW.termAssert') - if break_on_terminator_drop or break_every_terminator or break_every_step: - cut_point_rules.append('KMIR-CONTROL-FLOW.termDrop') - if break_on_terminator_unreachable or break_every_terminator or break_every_step: - cut_point_rules.append('KMIR-CONTROL-FLOW.termUnreachable') - return cut_point_rules - @staticmethod def from_kompiled_kore( smir_info: SMIRInfo, @@ -183,126 +118,11 @@ def run_smir( result = self.run_pattern(init_kore, depth=depth) return result - def apr_proof_from_smir( - self, - id: str, - smir_info: SMIRInfo, - start_symbol: str = 'main', - proof_dir: Path | None = None, - ) -> APRProof: - lhs_config, constraints = make_call_config( - self.definition, - smir_info=smir_info, - start_symbol=start_symbol, - mode=SymbolicMode(), - ) - lhs = CTerm(lhs_config, constraints) - - var_config, var_subst = split_config_from(lhs_config) - _rhs_subst: dict[str, KInner] = { - v_name: abstract_term_safely(KVariable('_'), base_name=v_name) for v_name in var_subst - } - _rhs_subst['K_CELL'] = KSequence([KMIR.Symbols.END_PROGRAM]) - rhs = CTerm(Subst(_rhs_subst)(var_config)) - kcfg = KCFG() - init_node = kcfg.create_node(lhs) - target_node = kcfg.create_node(rhs) - return APRProof(id, kcfg, [], init_node.id, target_node.id, {}, proof_dir=proof_dir) - @staticmethod def prove_rs(opts: ProveRSOpts) -> APRProof: - if not opts.rs_file.is_file(): - raise ValueError(f'Input file does not exist: {opts.rs_file}') - - label = str(opts.rs_file.stem) + '.' + opts.start_symbol - - with tempfile.TemporaryDirectory() as tmp_dir: - target_path = opts.proof_dir / label if opts.proof_dir is not None else Path(tmp_dir) - - if not opts.reload and opts.proof_dir is not None and APRProof.proof_data_exists(label, opts.proof_dir): - _LOGGER.info(f'Reading proof from disc: {opts.proof_dir}, {label}') - apr_proof = APRProof.read_proof_data(opts.proof_dir, label) - - smir_info = SMIRInfo.from_file(target_path / 'smir.json') - kmir = KMIR.from_kompiled_kore( - smir_info, - target_dir=target_path, - extra_module=opts.add_module, - bug_report=opts.bug_report, - symbolic=True, - haskell_target=opts.haskell_target, - llvm_lib_target=opts.llvm_lib_target, - ) - else: - _LOGGER.info(f'Constructing initial proof: {label}') - if opts.smir: - smir_info = SMIRInfo.from_file(opts.rs_file) - else: - smir_info = SMIRInfo(cargo_get_smir_json(opts.rs_file, save_smir=opts.save_smir)) - - smir_info = smir_info.reduce_to(opts.start_symbol) - # Report whether the reduced call graph includes any functions without MIR bodies - missing_body_syms = [ - sym - for sym, item in smir_info.items.items() - if 'MonoItemFn' in item['mono_item_kind'] - and item['mono_item_kind']['MonoItemFn'].get('body') is None - ] - has_missing = len(missing_body_syms) > 0 - _LOGGER.info(f'Reduced items table size {len(smir_info.items)}') - if has_missing: - _LOGGER.info(f'missing-bodies-present={has_missing} count={len(missing_body_syms)}') - _LOGGER.debug(f'Missing-body function symbols (first 5): {missing_body_syms[:5]}') - - kmir = KMIR.from_kompiled_kore( - smir_info, - target_dir=target_path, - extra_module=opts.add_module, - bug_report=opts.bug_report, - symbolic=True, - haskell_target=opts.haskell_target, - llvm_lib_target=opts.llvm_lib_target, - ) - - apr_proof = kmir.apr_proof_from_smir( - label, smir_info, start_symbol=opts.start_symbol, proof_dir=opts.proof_dir - ) - if apr_proof.proof_dir is not None and (apr_proof.proof_dir / label).is_dir(): - smir_info.dump(apr_proof.proof_dir / apr_proof.id / 'smir.json') - - if apr_proof.passed: - return apr_proof - - cut_point_rules = KMIR.cut_point_rules( - break_on_calls=opts.break_on_calls, - break_on_function_calls=opts.break_on_function_calls, - break_on_intrinsic_calls=opts.break_on_intrinsic_calls, - break_on_thunk=opts.break_on_thunk or opts.terminate_on_thunk, # must break for terminal rule - break_every_statement=opts.break_every_statement, - break_on_terminator_goto=opts.break_on_terminator_goto, - break_on_terminator_switch_int=opts.break_on_terminator_switch_int, - break_on_terminator_return=opts.break_on_terminator_return, - break_on_terminator_call=opts.break_on_terminator_call, - break_on_terminator_assert=opts.break_on_terminator_assert, - break_on_terminator_drop=opts.break_on_terminator_drop, - break_on_terminator_unreachable=opts.break_on_terminator_unreachable, - break_every_terminator=opts.break_every_terminator, - break_every_step=opts.break_every_step, - ) - - with kmir.kcfg_explore(label, terminate_on_thunk=opts.terminate_on_thunk) as kcfg_explore: - prover = APRProver( - kcfg_explore, - execute_depth=opts.max_depth, - cut_point_rules=cut_point_rules, - ) - prover.advance_proof( - apr_proof, - max_iterations=opts.max_iterations, - fail_fast=opts.fail_fast, - maintenance_rate=opts.maintenance_rate, - ) - return apr_proof + from ._prove import prove_rs + + return prove_rs(opts) class KMIRSemantics(DefaultSemantics):