Skip to content
Open
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
25 changes: 24 additions & 1 deletion src/kontrol/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import logging
import sys
from collections.abc import Iterable
from copy import copy
from typing import TYPE_CHECKING

from pyk.cli.pyk import parse_toml_args
Expand Down Expand Up @@ -48,6 +49,7 @@
)

if TYPE_CHECKING:
from argparse import Namespace
from pathlib import Path
from typing import Final, TypeVar

Expand Down Expand Up @@ -80,6 +82,13 @@

_LOGGER: Final = logging.getLogger(__name__)

_TOML_COMMAND_FALLBACKS: Final = {
'simplify-node': 'prove',
'step-node': 'prove',
'section-edge': 'prove',
'get-model': 'prove',
}


def _load_foundry(
foundry_root: Path,
Expand Down Expand Up @@ -110,7 +119,7 @@ def main() -> None:
parser = _create_argument_parser()
args = parser.parse_args()
args.config_file = config_file_path(args)
toml_args = parse_toml_args(args, get_option_string_destination, get_argument_type_setter)
toml_args = _parse_toml_args(args)
logging.basicConfig(
level=loglevel(args, toml_args),
format=_LOG_FORMAT,
Expand All @@ -131,6 +140,20 @@ def main() -> None:
execute(options)


def _parse_toml_args(args: Namespace) -> dict[str, object]:
command_toml_args = parse_toml_args(args, get_option_string_destination, get_argument_type_setter)

fallback_command = _TOML_COMMAND_FALLBACKS.get(args.command)
if fallback_command is None:
return command_toml_args

fallback_args = copy(args)
fallback_args.command = fallback_command
fallback_toml_args = parse_toml_args(fallback_args, get_option_string_destination, get_argument_type_setter)

return fallback_toml_args | command_toml_args


# Command implementation


Expand Down
118 changes: 118 additions & 0 deletions src/tests/unit/test_toml_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from pyk.cli.pyk import parse_toml_args
from pyk.cterm.symbolic import HASKELL_LOGGING_ENTRIES

from kontrol.__main__ import _parse_toml_args
from kontrol.cli import (
_create_argument_parser,
generate_options,
Expand Down Expand Up @@ -110,6 +111,123 @@ def test_toml_profiles() -> None:
assert args_dict['smt_timeout'] == 1000


def test_rpc_commands_fall_back_to_prove_profile(tmp_path: Path) -> None:
parser = _create_argument_parser()
toml_path = tmp_path / 'kontrol.toml'
toml_path.write_text(
'\n'.join(
[
'[prove.default]',
"foundry-project-root = '.'",
'workers = 4',
'smt-timeout = 1000',
'reinit = false',
]
)
)

def _args_dict(*cmd_args: str) -> dict:
args = parser.parse_args(list(cmd_args))
return _parse_toml_args(args)

simplify_node_args = _args_dict('simplify-node', '--config-file', str(toml_path), 'some_test', '1')
assert simplify_node_args['workers'] == 4
assert simplify_node_args['smt_timeout'] == 1000
assert simplify_node_args['reinit'] is False

step_node_args = _args_dict('step-node', '--config-file', str(toml_path), 'some_test', '1')
assert step_node_args['workers'] == 4
assert step_node_args['smt_timeout'] == 1000

section_edge_args = _args_dict('section-edge', '--config-file', str(toml_path), 'some_test', '1,2')
assert section_edge_args['workers'] == 4
assert section_edge_args['smt_timeout'] == 1000

get_model_args = _args_dict('get-model', '--config-file', str(toml_path), 'some_test')
assert get_model_args['workers'] == 4
assert get_model_args['smt_timeout'] == 1000


def test_rpc_commands_honor_selected_prove_profile(tmp_path: Path) -> None:
parser = _create_argument_parser()
toml_path = tmp_path / 'kontrol.toml'
toml_path.write_text(
'\n'.join(
[
'[prove.default]',
"foundry-project-root = '.'",
'workers = 1',
'smt-timeout = 250',
'',
'[prove.b_profile]',
"foundry-project-root = '.'",
'workers = 5',
'smt-timeout = 1000',
'reinit = true',
]
)
)

def _args_dict(*cmd_args: str) -> dict:
args = parser.parse_args(list(cmd_args))
return _parse_toml_args(args)

simplify_node_args = _args_dict(
'simplify-node', '--config-file', str(toml_path), '--config-profile', 'b_profile', 'some_test', '1'
)
assert simplify_node_args['workers'] == 5
assert simplify_node_args['smt_timeout'] == 1000
assert simplify_node_args['reinit'] is True

step_node_args = _args_dict(
'step-node', '--config-file', str(toml_path), '--config-profile', 'b_profile', 'some_test', '1'
)
assert step_node_args['workers'] == 5
assert step_node_args['smt_timeout'] == 1000
assert step_node_args['reinit'] is True

section_edge_args = _args_dict(
'section-edge', '--config-file', str(toml_path), '--config-profile', 'b_profile', 'some_test', '1,2'
)
assert section_edge_args['workers'] == 5
assert section_edge_args['smt_timeout'] == 1000
assert section_edge_args['reinit'] is True

get_model_args = _args_dict(
'get-model', '--config-file', str(toml_path), '--config-profile', 'b_profile', 'some_test'
)
assert get_model_args['workers'] == 5
assert get_model_args['smt_timeout'] == 1000
assert get_model_args['reinit'] is True


def test_command_profile_overrides_prove_fallback(tmp_path: Path) -> None:
parser = _create_argument_parser()
toml_path = tmp_path / 'kontrol.toml'
toml_path.write_text(
'\n'.join(
[
'[prove.default]',
"foundry-project-root = '.'",
'workers = 4',
'smt-timeout = 1000',
'reinit = false',
'',
'[simplify-node.default]',
'smt-timeout = 250',
'reinit = true',
]
)
)

args = parser.parse_args(['simplify-node', '--config-file', str(toml_path), 'some_test', '1'])
args_dict = _parse_toml_args(args)

assert args_dict['workers'] == 4
assert args_dict['smt_timeout'] == 250
assert args_dict['reinit'] is True


def _prove_options(cmd_args: list[str]) -> ProveOptions:
parser = _create_argument_parser()
args = parser.parse_args(['prove', *cmd_args])
Expand Down