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
@dataclass(frozen=True)
class OutputCtx:
type: str
package_prefix: str = ""
[docs]
@dataclass(frozen=True)
class RenderCtx:
env: jinja2.Environment
config: Any
output: OutputCtx
target: str
language: 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]"]
ALL_OUTPUT_TYPES = "*"
[docs]
@dataclass
class RenderRegistry:
_entries: dict[str, dict[type, _RendererFn]] = field(default_factory=dict)
def renders(
self,
render_type: type,
*,
output_type: str,
) -> Callable[[_RendererFn], _RendererFn]:
def decorator(fn: _RendererFn) -> _RendererFn:
self._entries.setdefault(output_type, dict())[render_type] = fn
return fn
return decorator
def render(
self,
obj: object,
ctx: RenderCtx,
) -> list[Fragment]:
render_type = type(obj)
fn = self._entries.get(ctx.output.type, {}).get(render_type)
fn = fn or self._entries.get(ALL_OUTPUT_TYPES, {}).get(render_type)
if fn:
return list(fn(obj, ctx))
return []
registry = RenderRegistry()