This subclass of `argparse.ArgumentParser` uses type hints on dataclasses to generate arguments. The class is designed to play well with the native argparse. In particular, you can add more (non-dataclass backed) arguments to the parser after initialization and you'll get the output ba
| 108 | |
| 109 | |
| 110 | class HfArgumentParser(ArgumentParser): |
| 111 | """ |
| 112 | This subclass of `argparse.ArgumentParser` uses type hints on dataclasses to generate arguments. |
| 113 | |
| 114 | The class is designed to play well with the native argparse. In particular, you can add more (non-dataclass backed) |
| 115 | arguments to the parser after initialization and you'll get the output back after parsing as an additional |
| 116 | namespace. Optional: To create sub argument groups use the `_argument_group_name` attribute in the dataclass. |
| 117 | """ |
| 118 | |
| 119 | dataclass_types: Iterable[DataClassType] |
| 120 | |
| 121 | def __init__(self, dataclass_types: Union[DataClassType, Iterable[DataClassType]], **kwargs): |
| 122 | """ |
| 123 | Args: |
| 124 | dataclass_types: |
| 125 | Dataclass type, or list of dataclass types for which we will "fill" instances with the parsed args. |
| 126 | kwargs (`Dict[str, Any]`, *optional*): |
| 127 | Passed to `argparse.ArgumentParser()` in the regular way. |
| 128 | """ |
| 129 | # To make the default appear when using --help |
| 130 | if "formatter_class" not in kwargs: |
| 131 | kwargs["formatter_class"] = ArgumentDefaultsHelpFormatter |
| 132 | super().__init__(**kwargs) |
| 133 | if dataclasses.is_dataclass(dataclass_types): |
| 134 | dataclass_types = [dataclass_types] |
| 135 | self.dataclass_types = list(dataclass_types) |
| 136 | for dtype in self.dataclass_types: |
| 137 | self._add_dataclass_arguments(dtype) |
| 138 | |
| 139 | @staticmethod |
| 140 | def _parse_dataclass_field(parser: ArgumentParser, field: dataclasses.Field): |
| 141 | field_name = f"--{field.name}" |
| 142 | kwargs = field.metadata.copy() |
| 143 | # field.metadata is not used at all by Data Classes, |
| 144 | # it is provided as a third-party extension mechanism. |
| 145 | if isinstance(field.type, str): |
| 146 | raise RuntimeError( |
| 147 | "Unresolved type detected, which should have been done with the help of " |
| 148 | "`typing.get_type_hints` method by default" |
| 149 | ) |
| 150 | |
| 151 | aliases = kwargs.pop("aliases", []) |
| 152 | if isinstance(aliases, str): |
| 153 | aliases = [aliases] |
| 154 | |
| 155 | origin_type = getattr(field.type, "__origin__", field.type) |
| 156 | if origin_type is Union or (hasattr(types, "UnionType") and isinstance(origin_type, types.UnionType)): |
| 157 | if str not in field.type.__args__ and ( |
| 158 | len(field.type.__args__) != 2 or type(None) not in field.type.__args__ |
| 159 | ): |
| 160 | raise ValueError( |
| 161 | "Only `Union[X, NoneType]` (i.e., `Optional[X]`) is allowed for `Union` because" |
| 162 | " the argument parser only supports one type per argument." |
| 163 | f" Problem encountered in field '{field.name}'." |
| 164 | ) |
| 165 | if type(None) not in field.type.__args__: |
| 166 | # filter `str` in Union |
| 167 | field.type = field.type.__args__[0] if field.type.__args__[1] is str else field.type.__args__[1] |
no outgoing calls