import importlib.metadata
from dataclasses import dataclass, field
from graphlib import CycleError, TopologicalSorter
from typing import TYPE_CHECKING, Any, NamedTuple
from pydantic import BaseModel
if TYPE_CHECKING:
from codegen.target import Target
_META_ATTR = "_operation_meta"
[docs]
class EmptyOptions(BaseModel):
pass
@dataclass(frozen=True)
class Propagate:
# indicates that options can be set at any level, not just at the scope
# specified in this operation meta
model: type[BaseModel]
@dataclass(frozen=True)
class Trigger:
"""What makes an operation run: a config scope, or an IR type.
An IR-triggered operation runs once per IR of that type as it is
emitted, wherever in the walk that happens, instead of once per
instance of a scope.
"""
scope: str | None = None
ir: type | None = None
def __post_init__(self) -> None:
if (self.scope is None) == (self.ir is None):
msg = "a trigger is either a scope name or an IR type, not both"
raise ValueError(msg)
@property
def is_ir(self) -> bool:
return bool(self.ir)
class OperationEntry(NamedTuple):
meta: OperationMeta
cls: type
[docs]
@dataclass
class OperationRegistry:
target: str
entries: list[OperationEntry] = field(default_factory=list)
def __init__(
self,
*operations: type,
target: str,
) -> None:
self.entries = []
self.target = target
self.register(*operations)
def register(self, *operations: type) -> None:
for cls in operations:
meta = getattr(cls, _META_ATTR, None)
if meta is None:
msg = f"{cls!r} is not an @operation-decorated class"
raise TypeError(msg)
if self.target in meta.targets:
self.entries.append(OperationEntry(meta=meta, cls=cls))
def validate_scopes(self, known: set[str]) -> None:
for entry in self.entries:
scope = entry.meta.trigger.scope
if scope is not None and scope not in known:
msg = (
f"Operation '{entry.meta.name}' targets "
f"scope '{scope}' which was not "
f"discovered from the config"
)
raise ValueError(msg)
def operations_for_ir(self, ir: object) -> list[OperationEntry]:
return _topo_sort(
[
entry
for entry in self.entries
if entry.meta.trigger.ir is type(ir)
]
)
def sorted_by_scope(self) -> dict[str, list[OperationEntry]]:
buckets: dict[str, list[OperationEntry]] = {}
for entry in self.entries:
if scope := entry.meta.trigger.scope:
buckets.setdefault(scope, []).append(entry)
return {name: _topo_sort(ops) for name, ops in buckets.items()}
[docs]
def operation(
name: str,
*,
trigger: Trigger,
targets: list[str],
requires: list[str] | None = None,
after_children: bool = False,
dispatch_on: str | None = None,
registry: OperationRegistry | None = None,
match_value: str | None = None,
options: type[BaseModel] | Propagate = EmptyOptions,
) -> Any: # noqa: ANN401
reqs: tuple[str, ...] = tuple(requires or [])
if trigger.is_ir and (dispatch_on or after_children):
msg = f"Operation '{name}' cannot use dispatch_on or after_children"
raise ValueError(msg)
meta = OperationMeta(
name=name,
trigger=trigger,
targets=targets,
requires=reqs,
after_children=after_children,
dispatch_on=dispatch_on,
match_value=match_value,
options_cls=(
options.model if isinstance(options, Propagate) else options
),
propagate_options=isinstance(options, Propagate),
)
def decorator(cls: type) -> type:
setattr(cls, _META_ATTR, meta)
if registry is not None:
registry.register(cls)
return cls
return decorator
[docs]
def load_registry(
target: Target,
plugin_operations: list[type] | None = None,
) -> OperationRegistry:
loaded_operations = [
entry_point.load()
for entry_point in importlib.metadata.entry_points(group="operations")
]
return OperationRegistry(
*loaded_operations,
*(plugin_operations or []),
target=target.name,
)
def _topo_sort(entries: list[OperationEntry]) -> list[OperationEntry]:
ordered = sorted(entries, key=lambda entry: entry.meta.name)
graph: dict[str, set[str]] = {
entry.meta.name: set(entry.meta.requires) for entry in ordered
}
_validate_requires(graph)
by_name: dict[str, OperationEntry] = {
entry.meta.name: entry for entry in ordered
}
try:
return [
by_name[name] for name in TopologicalSorter(graph).static_order()
]
except CycleError as exc:
msg = "Cycle detected in operation dependencies"
raise ValueError(msg) from exc
def _validate_requires(graph: dict[str, set[str]]) -> None:
for name, deps in graph.items():
for dependency in deps:
if dependency not in graph:
msg = (
f"Operation '{name}' requires '{dependency}', "
f"which is not registered"
)
raise ValueError(msg)