Source code for codegen.render

from collections.abc import Callable, Iterable
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Literal

from codegen.imports import ImportCollector
from codegen.store import BuildStore, NodePath

if TYPE_CHECKING:
    import jinja2


[docs] @dataclass(frozen=True) class RenderCtx: env: jinja2.Environment config: Any package_prefix: str = "" language: str = "" target_name: str = "" store: BuildStore = field(default_factory=BuildStore) instance_id: NodePath = field(default_factory=NodePath)
[docs] @dataclass class FileFragment: path: str template: str | None = None context: dict[str, Any] = field(default_factory=dict) imports: ImportCollector = field(default_factory=ImportCollector) if_exists: Literal["overwrite", "skip"] = "overwrite" executable: bool = False banner: bool = True id: str = "" def __post_init__(self) -> None: if not self.id: self.id = self.path def __or__(self, other: FileFragment) -> FileFragment: if self.template != other.template: msg = ( f"FileFragment template mismatch at {self.path!r}: " f"{self.template!r} vs {other.template!r}" ) raise ValueError(msg) if self.path != other.path: msg = ( f"FileFragment id {self.id!r} maps to two output " f"paths: {self.path!r} vs {other.path!r}" ) raise ValueError(msg) for key in self.context.keys() & other.context.keys(): if self.context[key] != other.context[key]: msg = ( f"FileFragment context conflict at {self.path!r} " f"for {key!r}: '{self.context[key]!r}' vs " f"'{other.context[key]!r}'" ) raise ValueError(msg) merged_if_exists: Literal["overwrite", "skip"] = ( "overwrite" if "overwrite" in (self.if_exists, other.if_exists) else "skip" ) return FileFragment( path=self.path, template=self.template, context=self.context | other.context, imports=self.imports | other.imports, if_exists=merged_if_exists, executable=self.executable or other.executable, banner=self.banner and other.banner, id=self.id, )
[docs] @dataclass class SnippetFragment: parent: str slot: str template: str | None = None context: dict[str, Any] = field(default_factory=dict) value: Any = None imports: ImportCollector = field(default_factory=ImportCollector) id: str | None = None def __or__(self, other: SnippetFragment) -> SnippetFragment: for attr in ("parent", "slot", "template", "value"): mine, theirs = getattr(self, attr), getattr(other, attr) if mine != theirs: msg = ( f"SnippetFragment id {self.id!r} {attr} mismatch: " f"{mine!r} vs {theirs!r}" ) raise ValueError(msg) for key in self.context.keys() & other.context.keys(): if self.context[key] != other.context[key]: msg = ( f"SnippetFragment id {self.id!r} context conflict " f"for {key!r}: {self.context[key]!r} vs " f"{other.context[key]!r}" ) raise ValueError(msg) return SnippetFragment( parent=self.parent, slot=self.slot, template=self.template, context=self.context | other.context, value=self.value, imports=self.imports | other.imports, id=self.id, )
type Fragment = FileFragment | SnippetFragment type _RendererFn = Callable[[Any, RenderCtx], "Iterable[Fragment]"]
[docs] @dataclass class RenderRegistry: _entries: dict[type, _RendererFn] = field(default_factory=dict)
[docs] def renders( self, output_type: type, ) -> Callable[[_RendererFn], _RendererFn]: """Register a renderer for *output_type*. Args: output_type: The output class this renderer handles. Returns: The original function, unmodified. """ def decorator(fn: _RendererFn) -> _RendererFn: self._entries[output_type] = fn return fn return decorator
def render( self, obj: object, ctx: RenderCtx, ) -> list[Fragment]: output_type = type(obj) fn = self._entries.get(output_type) if fn is None: msg = f"No renderer for {output_type.__name__}" raise LookupError(msg) return list(fn(obj, ctx))
registry = RenderRegistry()