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()