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
[docs] def format_imports(collector: ImportCollector, language: str) -> str: if not language: return "" return _get_formatter(language=language)(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)