from __future__ import annotations
from typing import Any
import flowrep as fr
import semantikon
from pyiron_snippets import retrieve
from pyiron_workflow import datatypes, execution
[docs]
class Atomic(datatypes.StaticNode[fr.schemas.AtomicRecipe, fr.schemas.AtomicData]):
def __init__(
self,
recipe: fr.schemas.AtomicRecipe,
label: fr.schemas.Label | None = None,
/,
**connections: datatypes.Port | datatypes.Node | fr.schemas.JSONABLE,
):
super().__init__(recipe, label, **connections)
func = retrieve.import_from_string(recipe.fully_qualified_name)
self._function_metadata = getattr(func, "_semantikon_metadata", None)
@classmethod
def _result_type(cls) -> type[fr.schemas.AtomicData]:
return fr.schemas.AtomicData
[docs]
def evaluate(
self,
run: execution.Run[execution.ResultType],
config: execution.RunConfig,
) -> execution.Run[execution.ResultType]:
output = _call_atomic(run.result)
_store_atomic_outputs(run.result, output)
return run
@property
def function_metadata(self) -> semantikon.FunctionMetadata | None:
return self._function_metadata
def _call_atomic(node: fr.schemas.AtomicData) -> Any:
"""
Invoke the underlying function, respecting positional-only parameter kinds.
Values are drawn from the live input ports; if a port has no value, its
default is used. A :class:`ValueError` is raised when neither is available.
"""
recipe = node.recipe
positional: list[Any] = []
keyword: dict[str, Any] = {}
for name in recipe.inputs:
port = node.input_ports[name]
val = (
port.value
if not isinstance(port.value, fr.schemas.NotData)
else port.default
)
if isinstance(val, fr.schemas.NotData):
raise ValueError(f"Input port '{name}' has no value and no default")
kind = recipe.reference.restricted_input_kinds.get(name)
if kind == fr.schemas.RestrictedParamKind.POSITIONAL_ONLY:
positional.append(val)
else:
keyword[name] = val
return node.function(*positional, **keyword)
def _store_atomic_outputs(node: fr.schemas.AtomicData, result: Any) -> None:
output_names = list(node.output_ports.keys())
if len(output_names) == 1:
node.output_ports[output_names[0]].value = result
else:
for name, val in zip(output_names, result, strict=True):
node.output_ports[name].value = val