lib.App: Refactor App / Cmd to adress multiple functionality and style issues #60
2 changed files with 140 additions and 76 deletions
|
|
@ -4,15 +4,18 @@ import asyncio
|
|||
import cProfile
|
||||
import os
|
||||
import sys
|
||||
import warnings
|
||||
|
||||
from argparse import ArgumentDefaultsHelpFormatter, ArgumentParser, Namespace
|
||||
from typing import TYPE_CHECKING, Any, cast, override
|
||||
from typing import TYPE_CHECKING, Any, NamedTuple, cast, override
|
||||
|
||||
from .AsyncRunner import AsyncRunner
|
||||
from .Cmd import AbstractCmd
|
||||
from .log import (
|
||||
DEBUG,
|
||||
ERR,
|
||||
NOTICE,
|
||||
WARNING,
|
||||
LogFlag,
|
||||
log,
|
||||
log_m,
|
||||
|
|
@ -31,25 +34,47 @@ if TYPE_CHECKING:
|
|||
from typing import TypeVar
|
||||
T = TypeVar('T')
|
||||
|
||||
def _get_current_event_loop() -> asyncio.AbstractEventLoop | None:
|
||||
"""Return the current event loop of this thread, or None if there is
|
||||
none, without creating one implicitly or emitting a deprecation
|
||||
warning."""
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter('error', DeprecationWarning)
|
||||
try:
|
||||
return asyncio.get_event_loop()
|
||||
except (RuntimeError, DeprecationWarning):
|
||||
return None
|
||||
|
||||
class _SubCommand(NamedTuple):
|
||||
|
||||
cmd: AbstractCmd
|
||||
parser: ArgumentParser
|
||||
|
||||
class App: # export
|
||||
|
||||
def _add_arguments(self, parser: ArgumentParser) -> None:
|
||||
self.__parser.add_argument(
|
||||
'--log-flags', help = 'Log flags', default = self.__default_log_flags
|
||||
parser.add_argument(
|
||||
'--log-flags',
|
||||
help = 'Log flags',
|
||||
default = self.__default_log_flags,
|
||||
type = parse_log_flags,
|
||||
)
|
||||
self.__parser.add_argument(
|
||||
'--log-level', help = 'Log level', default = self.__default_log_level
|
||||
parser.add_argument(
|
||||
'--log-level',
|
||||
help = 'Log level',
|
||||
default = self.__default_log_level,
|
||||
type = parse_log_level,
|
||||
)
|
||||
self.__parser.add_argument(
|
||||
parser.add_argument(
|
||||
'--log-file', help = 'Log file', default = self.__default_log_file
|
||||
)
|
||||
self.__parser.add_argument(
|
||||
parser.add_argument(
|
||||
'--backtrace',
|
||||
help = 'Show exception backtraces',
|
||||
action = 'store_true',
|
||||
default = self.__back_trace,
|
||||
)
|
||||
self.__parser.add_argument(
|
||||
parser.add_argument(
|
||||
'--write-profile',
|
||||
help = 'Profile code and store output to file',
|
||||
default = None,
|
||||
|
|
@ -87,6 +112,45 @@ class App: # export
|
|||
eloop: asyncio.AbstractEventLoop | None = None,
|
||||
) -> None:
|
||||
|
||||
self.__args: Namespace | None = None
|
||||
self.__cmdline: str | None = None
|
||||
self.__description = description
|
||||
|
||||
self.__default_log_flags = self._default_log_flags(
|
||||
LogFlag.STDERR | LogFlag.POSITION | LogFlag.PRIO | LogFlag.COLOR
|
||||
)
|
||||
if (env := os.getenv(self._default_log_flags_env(), None)) is not None:
|
||||
self.__default_log_flags = parse_log_flags(env)
|
||||
|
||||
self.__default_log_level = self._default_log_level(NOTICE)
|
||||
if (env := os.getenv(self._default_log_level_env(), None)) is not None:
|
||||
self.__default_log_level = parse_log_level(env)
|
||||
|
||||
self.__default_log_file = self._default_log_file(None)
|
||||
if (env := os.getenv(self._default_log_file_env(), None)) is not None:
|
||||
self.__default_log_file = env
|
||||
|
||||
self.__back_trace = self._default_show_backtrace(False)
|
||||
if (env := os.getenv(self._default_show_backtrace_env(), None)) is not None:
|
||||
self.__back_trace = env.lower() in ['1', 'true']
|
||||
|
||||
set_log_flags(self.__default_log_flags)
|
||||
set_log_level(self.__default_log_level)
|
||||
|
||||
self.__async_runner: AsyncRunner | None = None
|
||||
self.__eloop = eloop
|
||||
self.__own_eloop = False
|
||||
|
||||
cmd_classes: LoadTypes[AbstractCmd] = LoadTypes(
|
||||
modules if modules else ['__main__'],
|
||||
type_name_filter = name_filter,
|
||||
type_filter = [AbstractCmd],
|
||||
)
|
||||
self.__cmds: list[AbstractCmd] = [cmd_class(self) for cmd_class in cmd_classes]
|
||||
self._build_parser()
|
||||
|
||||
def _build_parser(self, argv: list[str] | None = None) -> None:
|
||||
|
||||
def add_cmd_to_parser(cmd: AbstractCmd, parsers: Any) -> ArgumentParser:
|
||||
parser = cast(
|
||||
'ArgumentParser',
|
||||
|
|
@ -112,23 +176,25 @@ class App: # export
|
|||
if not cmds:
|
||||
return
|
||||
|
||||
class SubCommand:
|
||||
|
||||
def __init__(self, cmd: AbstractCmd, parser: Any):
|
||||
self.cmd = cmd
|
||||
self.parser = parser
|
||||
|
||||
title = 'Available subcommands'
|
||||
if isinstance(parent, AbstractCmd):
|
||||
title += ' of ' + parent.name
|
||||
subparsers = parser.add_subparsers(
|
||||
title = title, metavar = '', dest = 'command'
|
||||
)
|
||||
scs: dict[str, SubCommand] = {}
|
||||
scs: dict[str, _SubCommand] = {}
|
||||
for cmd in cmds:
|
||||
cmd.set_parent(parent)
|
||||
scs[cmd.name] = SubCommand(cmd, add_cmd_to_parser(cmd, subparsers))
|
||||
if cmd.name in scs:
|
||||
log(WARNING, f'Duplicate subcommand name: {cmd.name}')
|
||||
scs[cmd.name] = _SubCommand(cmd, add_cmd_to_parser(cmd, subparsers))
|
||||
for alias in cmd.aliases:
|
||||
if alias != cmd.name and alias in scs:
|
||||
log(
|
||||
WARNING,
|
||||
f'Subcommand alias "{alias}" of "{cmd.name}" '
|
||||
'collides with an earlier subcommand',
|
||||
)
|
||||
scs[alias] = scs[cmd.name]
|
||||
if all:
|
||||
seen: set[int] = set()
|
||||
|
|
@ -139,69 +205,37 @@ class App: # export
|
|||
sc.cmd, sc.parser, sc.cmd.children, all = all
|
||||
)
|
||||
return
|
||||
args, _ = self.__parser.parse_known_args()
|
||||
# -- Re-parse the command line to find the invoked subcommand.
|
||||
# This works because every level below uses dest = 'command',
|
||||
# so each pass descends one level further into the command
|
||||
# tree.
|
||||
args, _ = self.__parser.parse_known_args(argv)
|
||||
cmd_name = getattr(args, 'command', None)
|
||||
if cmd_name in scs:
|
||||
sc = scs[cmd_name]
|
||||
add_cmds_to_parser(sc.cmd, sc.parser, sc.cmd.children, all = all)
|
||||
|
||||
from .Cmd import AbstractCmd
|
||||
|
||||
self.__args: Namespace | None = None
|
||||
self.__cmdline: str | None = None
|
||||
|
||||
self.__default_log_flags = self._default_log_flags(
|
||||
LogFlag.STDERR | LogFlag.POSITION | LogFlag.PRIO | LogFlag.COLOR
|
||||
cmdline = sys.argv if argv is None else argv
|
||||
if argv is None:
|
||||
argv = sys.argv[1:]
|
||||
add_all_parsers = (
|
||||
'-h' in argv or '--help' in argv or '_ARGCOMPLETE' in os.environ
|
||||
)
|
||||
if (env := os.getenv(self._default_log_flags_env(), None)) is not None:
|
||||
self.__default_log_flags = parse_log_flags(env)
|
||||
|
||||
self.__default_log_level = self._default_log_level(NOTICE)
|
||||
if (env := os.getenv(self._default_log_level_env(), None)) is not None:
|
||||
self.__default_log_level = parse_log_level(env)
|
||||
|
||||
self.__default_log_file = self._default_log_file(None)
|
||||
if (env := os.getenv(self._default_log_file_env(), None)) is not None:
|
||||
self.__default_log_file = env
|
||||
|
||||
self.__back_trace = self._default_show_backtrace(False)
|
||||
if (env := os.getenv(self._default_show_backtrace_env())) is not None:
|
||||
self.__back_trace = env.lower() in ['1', 'true']
|
||||
|
||||
set_log_flags(self.__default_log_flags)
|
||||
set_log_level(self.__default_log_level)
|
||||
|
||||
self.__async_runner: AsyncRunner | None = None
|
||||
self.__eloop = eloop
|
||||
self.__own_eloop = False
|
||||
|
||||
self.__parser = ArgumentParser(
|
||||
formatter_class = ArgumentDefaultsHelpFormatter,
|
||||
description = description,
|
||||
description = self.__description,
|
||||
add_help = False,
|
||||
)
|
||||
self._add_arguments(self.__parser)
|
||||
|
||||
args, _ = self.__parser.parse_known_args()
|
||||
args, _ = self.__parser.parse_known_args(argv)
|
||||
set_log_flags(args.log_flags)
|
||||
set_log_level(args.log_level)
|
||||
|
||||
log(DEBUG, f'-------------- Running: >{pretty_cmd(sys.argv)}<')
|
||||
log(DEBUG, f'-------------- Running: >{pretty_cmd(cmdline)}<')
|
||||
|
||||
cmd_classes: LoadTypes[AbstractCmd] = LoadTypes(
|
||||
modules if modules else ['__main__'],
|
||||
type_name_filter = name_filter,
|
||||
type_filter = [AbstractCmd],
|
||||
)
|
||||
add_all_parsers = (
|
||||
'-h' in sys.argv or '--help' in sys.argv or '_ARGCOMPLETE' in os.environ
|
||||
)
|
||||
add_cmds_to_parser(
|
||||
self,
|
||||
self.__parser,
|
||||
[cmd_class(self) for cmd_class in cmd_classes],
|
||||
all = add_all_parsers,
|
||||
)
|
||||
add_cmds_to_parser(self, self.__parser, self.__cmds, all = add_all_parsers)
|
||||
|
||||
# -- Add help only now, wouldn't want to have parse_known_args() exit
|
||||
# on --help with subcommands missing
|
||||
|
|
@ -210,14 +244,19 @@ class App: # export
|
|||
)
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the application and release all resources"""
|
||||
if self.__async_runner is not None:
|
||||
self.__async_runner.close()
|
||||
self.__async_runner = None
|
||||
if self.__own_eloop:
|
||||
if self.__eloop is not None:
|
||||
if not self.__eloop.is_closed():
|
||||
self.__eloop.close()
|
||||
self.__eloop = None
|
||||
self.__own_eloop = False
|
||||
|
||||
async def __aenter__(self) -> None:
|
||||
pass
|
||||
async def __aenter__(self) -> App:
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
|
|
@ -225,7 +264,7 @@ class App: # export
|
|||
exc: BaseException | None,
|
||||
tb: types.TracebackType | None,
|
||||
) -> None:
|
||||
pass
|
||||
self.close()
|
||||
|
||||
async def __run(self, argv: list[str] | None = None) -> None:
|
||||
|
||||
|
|
@ -245,8 +284,10 @@ class App: # export
|
|||
|
||||
argcomplete.autocomplete(self.__parser, default_completer = NoopCompleter())
|
||||
|
||||
except Exception:
|
||||
pass
|
||||
except ImportError:
|
||||
log(DEBUG, 'argcomplete is not installed, shell completion disabled')
|
||||
except Exception as e:
|
||||
log(DEBUG, f'Shell completion disabled: {e}')
|
||||
|
||||
self.__args = self.__parser.parse_args(args = argv)
|
||||
|
||||
|
|
@ -261,9 +302,17 @@ class App: # export
|
|||
pr.enable()
|
||||
|
||||
try:
|
||||
ret = await self._run(self.__args)
|
||||
if isinstance(ret, int) and ret >= 0 and ret <= 0xFF:
|
||||
exit_status = ret
|
||||
result = await self._run(self.__args)
|
||||
if isinstance(result, int):
|
||||
if 0 <= result <= 0xFF:
|
||||
exit_status = result
|
||||
else:
|
||||
log(
|
||||
WARNING,
|
||||
f'Command returned invalid exit status {result}, '
|
||||
'using 1 instead',
|
||||
)
|
||||
exit_status = 1
|
||||
except Exception as e:
|
||||
log_m(ERR, f'Failed: {repr(e) if self.__back_trace else str(e)}')
|
||||
exit_status = 1
|
||||
|
|
@ -288,11 +337,11 @@ class App: # export
|
|||
# want to do something else, for instance if you don't have sub-commands,
|
||||
# or if want to do anything before and / or after the subcommands.
|
||||
async def _run(self, args: Namespace) -> None | int:
|
||||
if not hasattr(self.__args, 'func'):
|
||||
if not hasattr(args, 'func'):
|
||||
self.__parser.print_help()
|
||||
return None
|
||||
# Run sub-command. Overwrite if you want to do anything before or after
|
||||
return cast('None | int', await self.args.func(args))
|
||||
return cast('None | int', await args.func(args))
|
||||
|
||||
def call_async(self, awaitable: Awaitable[T], timeout: float | None = None) -> T:
|
||||
return self.async_runner.call(awaitable, timeout)
|
||||
|
|
@ -330,18 +379,25 @@ class App: # export
|
|||
return self.__parser
|
||||
|
||||
def run(self, argv: list[str] | None = None) -> None:
|
||||
previous_eloop: asyncio.AbstractEventLoop | None = None
|
||||
if self.__eloop is None:
|
||||
previous_eloop = _get_current_event_loop()
|
||||
eloop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(eloop)
|
||||
self.__eloop = eloop
|
||||
self.__own_eloop = True
|
||||
try:
|
||||
if argv is not None:
|
||||
self._build_parser(argv)
|
||||
ret = self.eloop.run_until_complete(self.__run(argv))
|
||||
finally:
|
||||
if self.__async_runner:
|
||||
self.__async_runner.close()
|
||||
self.__async_runner = None
|
||||
self.close()
|
||||
# -- Restore the event loop the thread had before run(), or
|
||||
# unset the loop if there was none.
|
||||
if previous_eloop is not None:
|
||||
asyncio.set_event_loop(previous_eloop)
|
||||
else:
|
||||
asyncio.set_event_loop(None)
|
||||
return ret
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -88,7 +88,15 @@ class AbstractCmd(abc.ABC):
|
|||
self, cmds: Cmd | list[Cmd] | Types[Any] | list[Types[Any]]
|
||||
) -> None:
|
||||
if isinstance(cmds, Cmd):
|
||||
raise NotImplementedError('Single Cmd should be handled elsewhere')
|
||||
if any(child.name == cmds.name for child in self.__children):
|
||||
raise Exception(
|
||||
f'Can\'t register subcommand with already taken name "{cmds.name}"'
|
||||
)
|
||||
self.__child_classes.append(type(cmds))
|
||||
cmds.set_parent(self)
|
||||
self.__children.append(cmds)
|
||||
assert len(self.__children) == len(self.__child_classes)
|
||||
return
|
||||
if isinstance(cmds, list):
|
||||
for cmd in cmds:
|
||||
self.add_subcommands(cmd)
|
||||
|
|
|
|||
Loading…
Reference in a new issue