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
2 changes: 1 addition & 1 deletion .copier-answers.yaml
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# Changes here will be overwritten by Copier
_commit: ebb6e1c
_commit: '4372075'
_src_path: https://github.com/python-project-templates/base.git
add_docs: false
add_extension: python
Expand Down
2 changes: 1 addition & 1 deletion .github/workflows/build.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ on:
workflow_dispatch:

concurrency:
group: ${{ github.workflow }}-${{ github.head_ref || github.run_id }}
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true

permissions:
Expand Down
10 changes: 5 additions & 5 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -25,16 +25,16 @@ lint-py: ## lint python with ruff
python -m ruff format --check p2a

lint-docs: ## lint docs with mdformat and codespell
python -m mdformat --check README.md
python -m codespell_lib README.md
python -m mdformat --check README.md $(wildcard docs/src)
python -m codespell_lib README.md $(wildcard docs/src)

fix-py: ## autoformat python code with ruff
python -m ruff check --fix p2a
python -m ruff format p2a

fix-docs: ## autoformat docs with mdformat and codespell
python -m mdformat README.md
python -m codespell_lib --write README.md
python -m mdformat README.md $(wildcard docs/src)
python -m codespell_lib --write README.md $(wildcard docs/src)

lint: lint-py lint-docs ## run all linters
lints: lint
Expand All @@ -52,7 +52,7 @@ check-dist: ## check python sdist and wheel with check-dist
check-types: ## check python types with ty
ty check --python $$(which python)

checks: check-dist
checks: check-dist check-types

# Alias
check: checks
Expand Down
23 changes: 14 additions & 9 deletions p2a/model.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,10 @@
import sys
from argparse import ArgumentParser
from enum import Enum
from logging import Logger
from pathlib import Path
from types import UnionType
from typing import Literal, Union, get_args, get_origin
from typing import Any, Literal, TypeVar, Union, cast, get_args, get_origin

from pkn.logging import getSimpleLogger

Expand All @@ -18,7 +19,8 @@
"parse_extra_args_model",
)

_log = None
_log: Logger
ModelT = TypeVar("ModelT", bound=BaseModel)


def _initlog(level: str = "WARNING"):
Expand All @@ -38,7 +40,7 @@ def _add_argument(parser: ArgumentParser, name: str, arg_type: type, default_val
parser.add_argument(name, type=arg_type, default=default_value, **kwargs)


def parse_extra_args(subparser: ArgumentParser | None = None, argv: list[str] | None = None) -> list[str]:
def parse_extra_args(subparser: ArgumentParser | None = None, argv: list[str] | None = None) -> tuple[dict[str, Any], list[str]]:
if subparser is None:
subparser = ArgumentParser(prog="p2a", allow_abbrev=False)

Expand Down Expand Up @@ -126,7 +128,7 @@ def _recurse_add_fields(parser: ArgumentParser, model: Union["BaseModel", type["
########################
# MARK: str, int, float
try:
_add_argument(parser=parser, name=arg_name, arg_type=field_type, default_value=default_value)
_add_argument(parser=parser, name=arg_name, arg_type=cast(type, field_type), default_value=default_value)
except TypeError:
# TODO: handle more complex types if needed
_add_argument(parser=parser, name=arg_name, arg_type=str, default_value=default_value)
Expand Down Expand Up @@ -200,7 +202,10 @@ def _recurse_add_fields(parser: ArgumentParser, model: Union["BaseModel", type["
#################################
# MARK: List[str|int|float|bool]
_add_argument(
parser=parser, name=arg_name, arg_type=str, default_value=",".join(map(str, default_value)) if isinstance(field, str) else None
parser=parser,
name=arg_name,
arg_type=str,
default_value=",".join(map(str, default_value)) if isinstance(default_value, list) else None,
)
elif get_origin(field_type) in (dict, dict):
######################
Expand Down Expand Up @@ -319,7 +324,7 @@ def create_model_parser(model: "BaseModel") -> ArgumentParser:
return parser


def parse_extra_args_model(model: "BaseModel", argv: list[str] | None = None) -> Union["BaseModel", dict]:
def parse_extra_args_model(model: ModelT, argv: list[str] | None = None) -> tuple[ModelT, list[str]]:
# Parse the extra args and update the model
args, kwargs = parse_extra_args(create_model_parser(model), argv)

Expand All @@ -330,8 +335,8 @@ def parse_extra_args_model(model: "BaseModel", argv: list[str] | None = None) ->
parts = key.split(".")

# Accounting
sub_model = model
parent_model = None
sub_model: Any = model
parent_model: Any = None

for i, part in enumerate(parts[:-1]):
if part.isdigit() and isinstance(sub_model, list):
Expand Down Expand Up @@ -457,7 +462,7 @@ def parse_extra_args_model(model: "BaseModel", argv: list[str] | None = None) ->

# Grab the field from the model class and make a type adapter
field = model_to_set.__class__.model_fields[key]
adapter = TypeAdapter(field.annotation)
adapter = TypeAdapter(cast(Any, field.annotation))

if value is not None:
_log.debug(f"Setting field '{key}' on model '{model_to_set.__class__.__name__}' with raw value '{value}'")
Expand Down
12 changes: 10 additions & 2 deletions p2a/tests/test_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@

from pydantic import BaseModel

from p2a import parse_extra_args_model
from p2a import create_model_parser, parse_extra_args_model
from p2a.model import _initlog


Expand Down Expand Up @@ -46,7 +46,7 @@ class MyTopLevelModel(BaseModel, validate_assignment=True):
dict_enum_key_model_value: dict[MyEnum, SubModel] = {MyEnum.OPTION_A: SubModel()}

submodel: SubModel
submodel2: SubModel = SubModel(sub_args=84, sub_arg_with_value="predefined", sub_arg_enum=MyEnum.OPTION_B, sub_arg_literal="z")
submodel2: SubModel = SubModel(sub_arg=84, sub_arg_with_value="predefined", sub_arg_enum=MyEnum.OPTION_B, sub_arg_literal="z")
submodel3: SubModel | None = None

submodel_list_instanced: list[SubModel] = [SubModel()]
Expand All @@ -62,6 +62,13 @@ class MyTopLevelModel(BaseModel, validate_assignment=True):


class TestCLIMdel:
def test_list_default(self):
class Model(BaseModel):
values: list[int] = [1, 2, 3]

parser = create_model_parser(Model())
assert parser.parse_args([]).values == "1,2,3"

def test_get_arg_from_model(self):
with (
patch.object(
Expand Down Expand Up @@ -164,6 +171,7 @@ def test_get_arg_from_model(self):
assert model.submodel2.sub_arg_enum == MyEnum.OPTION_B
assert model.submodel2.sub_arg_literal == "z"

assert model.submodel3 is not None
assert model.submodel3.sub_arg == 300
assert model.submodel_list_instanced[0].sub_arg == 400
assert model.submodel_list_instanced[0].sub_arg_with_value == "list_value"
Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@ develop = [
"codespell",
"hatchling",
"mdformat",
"mdformat-frontmatter",
"mdformat-tables>=1",
"pytest",
"pytest-cov",
Expand Down Expand Up @@ -82,7 +83,7 @@ replace = 'version = "{new_version}"'
[tool.coverage.run]
branch = true
omit = [
"p2a/tests/integration/",
"p2a/tests/integration/*",
]

[tool.coverage.report]
Expand Down