Source code for codegen.engine

from collections import ChainMap
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, cast

from pydantic import BaseModel

from codegen.scope import PROJECT, Scope, ScopeTree, discover_scopes
from codegen.store import BuildStore, NodePath

if TYPE_CHECKING:
    from codegen.operation import (
        OperationEntry,
        OperationMeta,
        OperationRegistry,
    )


[docs] @dataclass class BuildContext[InstanceT, ConfigT: BaseModel]: config: ConfigT scope: Scope instance: InstanceT instance_id: NodePath store: BuildStore package_prefix: str = ""
[docs] class Engine: registry: OperationRegistry package_prefix: str = "" _config: BaseModel _store: BuildStore _ops: dict[str, list[OperationEntry]] _scope_tree: ScopeTree def __init__( self, registry: OperationRegistry, package_prefix: str = "", ) -> None: self.registry = registry self.package_prefix = package_prefix def build(self, config: BaseModel) -> BuildStore: self._scope_tree = discover_scopes(type(config)) self.registry.validate_scopes( {scope.name for scope in self._scope_tree} ) self._config = config self._store = BuildStore(scope_tree=self._scope_tree) self._ops = self.registry.sorted_by_scope() self._visit( scope=PROJECT, instance_config=config, instance_id=NodePath(("project",)), ) return self._store def _visit( self, scope: Scope, instance_config: BaseModel, instance_id: NodePath, ) -> None: self._store.register_instance( instance_id, instance_config, ) ctx = BuildContext( config=self._config, scope=scope, instance=instance_config, instance_id=instance_id, store=self._store, package_prefix=self.package_prefix, ) ops = self._ops.get(scope.name, []) self._run_ops(ops, ctx, after_children=False) for child_scope in self._scope_tree.children_of(scope): scope_configs = _resolve_scope_list( scope=child_scope, parent_config=instance_config, ) index = 0 while index < len(scope_configs): self._visit( scope=child_scope, instance_config=scope_configs[index], instance_id=NodePath.combine( instance_id, f"{child_scope.config_key}.{index}" ), ) index += 1 self._run_ops(ops, ctx, after_children=True) def _run_ops( self, ops: list[OperationEntry], ctx: BuildContext[Any, BaseModel], *, after_children: bool, ) -> None: for meta, op_cls in ops: if meta.after_children != after_children: continue # dispatch_on: fire only on the instance whose discriminator # matches this op's name (e.g. OperationConfig.name == "get"). if meta.dispatch_on is not None and getattr( ctx.instance, meta.dispatch_on, None ) != (meta.match_value or meta.name): continue self._run_operation(meta, op_cls, ctx) def _run_operation( self, meta: OperationMeta, op_cls: type, ctx: BuildContext[Any, BaseModel], *triggering_ir: object, ) -> None: operation_instance = op_cls() when_method = getattr(operation_instance, "when", None) options = _resolve_options(meta, ctx) if callable(when_method) and not when_method(ctx, options): return for new_ir in operation_instance.build(ctx, options, *triggering_ir): ctx.store.add(ctx.instance_id, meta.name, new_ir) for triggered in self.registry.operations_for_ir(new_ir): self._run_operation( triggered.meta, triggered.cls, ctx, new_ir, )
def _resolve_scope_list( *, scope: Scope, parent_config: BaseModel, ) -> list[BaseModel]: attr_value: object = parent_config for attr in scope.resolve_path: attr_value = getattr(attr_value, attr) return cast("list[BaseModel]", attr_value) def _resolve_options( meta: OperationMeta, ctx: BuildContext[Any, BaseModel], ) -> BaseModel: sources = [ctx.instance] if meta.propagate_options: sources.extend(ctx.store.ancestors(ctx.instance_id)) option_dicts = [ options for source in sources if isinstance(options := getattr(source, "options", None), dict) ] return meta.options_cls(**ChainMap(*option_dicts))