Source code for codegen.operation

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)
[docs] @dataclass(frozen=True) class OperationMeta: name: str trigger: Trigger targets: list[str] requires: tuple[str, ...] = () after_children: bool = False dispatch_on: str | None = None match_value: str | None = None options_cls: type[BaseModel] = EmptyOptions propagate_options: bool = False
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)