Source code for codegen.imports
import functools
import importlib.metadata
from dataclasses import dataclass
from typing import TYPE_CHECKING
from codegen.naming import Name
if TYPE_CHECKING:
from collections.abc import Callable
_ENTRY_POINT_GROUP = "codegen.import_formatters"
type NameOrNameWithAlias = Name | tuple[Name, str]
@dataclass(frozen=True, eq=True)
class _Import:
name: Name
alias: str | None
[docs]
class ImportCollector:
_imports: set[_Import]
def __init__(self, *imports: NameOrNameWithAlias) -> None:
self._imports = set()
for _import in imports:
if isinstance(_import, Name):
self.add_name(_import)
else:
_import, alias = _import
self.add_name(_import, alias=alias)
def add_from(
self,
module: str,
*names: str,
alias: str | None = None,
) -> None:
if not isinstance(module, str):
raise TypeError(f"import module must be str, got {module!r}")
for name in names:
if not isinstance(name, str):
msg = f"import name for {module!r} must be str, got {name!r}"
raise TypeError(msg)
constructed_names = {
_Import(Name.from_module_name(module, name).external, alias)
for name in names
}
self._imports |= constructed_names
def add_name(self, *names: Name, alias: str | None = None) -> None:
self._imports |= {_Import(name, alias) for name in names}
def update(self, other: ImportCollector) -> None:
self._imports |= other.imports
def __or__(self, other: ImportCollector) -> ImportCollector:
merged = ImportCollector()
merged.update(self)
merged.update(other)
return merged
@property
def imports(self) -> set[_Import]:
return self._imports
def consume(self) -> ImportCollector:
new_collector = ImportCollector()
new_collector.update(self)
self._imports = set()
return new_collector
@functools.cache
def _get_formatter(language: str) -> Callable[[ImportCollector], str]:
available = list(importlib.metadata.entry_points(group=_ENTRY_POINT_GROUP))
for entry_point in available:
if entry_point.name == language:
return entry_point.load()
registered = sorted(entry_point.name for entry_point in available)
msg = (
f"No import formatter registered for language {language!r}; "
f"registered: {registered}"
)
raise KeyError(msg)