1from __future__ import annotations
2
3import inspect
4from collections.abc import Callable
5from contextlib import chdir
6from dataclasses import dataclass
7from functools import wraps
8from tempfile import TemporaryDirectory
9from typing import Any, Generic, TypeVar
10
11from daggerml import Node
12from daggerml.api import DmlRepoError
13from daggerml.contrib.api import DelayedRunnable
14
15T = TypeVar("T")
16
17
18@dataclass(frozen=True)
19class MockNode(Generic[T]):
20 _value: T
21
22 def value(self) -> T:
23 return self._value
24
25 @classmethod
26 def from_value(cls, value: T) -> MockNode | Node:
27 if isinstance(value, (Node, MockNode)):
28 return value
29 return cls(value)
30
31
32def wrap_node(arg: Any) -> Any:
33 if isinstance(arg, (Node, MockNode)):
34 return arg
35 return MockNode(arg)
36
37
38def defunkify(value: DelayedRunnable) -> Callable[..., Any]:
39 current = value
40 while isinstance(current.sub, DelayedRunnable):
41 current = current.sub
42 if current.uri != "script":
43 raise DmlRepoError("defunkify requires innermost script delayed runnable")
44 fn = current.kwargs.get("fn")
45 if not callable(fn):
46 raise DmlRepoError("defunkify requires callable fn in innermost script kwargs")
47
48 sig = inspect.signature(fn)
49 param_names = tuple(sig.parameters)
50
51 @wraps(fn)
52 def wrapped(*args: Any, **kwargs: Any) -> Any:
53 bound = sig.bind_partial(*args, **kwargs)
54 bound.apply_defaults()
55 for name, param in sig.parameters.items():
56 if name not in bound.arguments or name == param_names[0]:
57 continue
58 if param.kind == inspect.Parameter.VAR_POSITIONAL:
59 bound.arguments[name] = tuple(wrap_node(arg) for arg in bound.arguments[name])
60 elif param.kind == inspect.Parameter.VAR_KEYWORD:
61 bound.arguments[name] = {key: wrap_node(arg) for key, arg in bound.arguments[name].items()}
62 else:
63 bound.arguments[name] = wrap_node(bound.arguments[name])
64 with TemporaryDirectory(prefix="dml-defunkify-") as tmpd:
65 with chdir(tmpd):
66 return fn(*bound.args, **bound.kwargs)
67
68 return wrapped
69
70
71__all__ = ["MockNode", "defunkify"]