daggerml.contrib.testing

View source
 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"]

MockNode

@dataclass(frozen=True)
class MockNode(typing.Generic[~T]):
View source
19@dataclass(frozen=True)
20class MockNode(Generic[T]):
21    _value: T
22
23    def value(self) -> T:
24        return self._value
25
26    @classmethod
27    def from_value(cls, value: T) -> MockNode | Node:
28        if isinstance(value, (Node, MockNode)):
29            return value
30        return cls(value)

MockNode.__init__

MockNode(_value: ~T)

MockNode.value

def value(self) -> ~T:
View source
23    def value(self) -> T:
24        return self._value

MockNode.from_value

@classmethod
def from_value(cls, value: ~T) -> MockNode | daggerml.Node:
View source
26    @classmethod
27    def from_value(cls, value: T) -> MockNode | Node:
28        if isinstance(value, (Node, MockNode)):
29            return value
30        return cls(value)

defunkify

def defunkify( value: daggerml.contrib.codecs.DelayedRunnable) -> Callable[..., typing.Any]:
View source
39def defunkify(value: DelayedRunnable) -> Callable[..., Any]:
40    current = value
41    while isinstance(current.sub, DelayedRunnable):
42        current = current.sub
43    if current.uri != "script":
44        raise DmlRepoError("defunkify requires innermost script delayed runnable")
45    fn = current.kwargs.get("fn")
46    if not callable(fn):
47        raise DmlRepoError("defunkify requires callable fn in innermost script kwargs")
48
49    sig = inspect.signature(fn)
50    param_names = tuple(sig.parameters)
51
52    @wraps(fn)
53    def wrapped(*args: Any, **kwargs: Any) -> Any:
54        bound = sig.bind_partial(*args, **kwargs)
55        bound.apply_defaults()
56        for name, param in sig.parameters.items():
57            if name not in bound.arguments or name == param_names[0]:
58                continue
59            if param.kind == inspect.Parameter.VAR_POSITIONAL:
60                bound.arguments[name] = tuple(wrap_node(arg) for arg in bound.arguments[name])
61            elif param.kind == inspect.Parameter.VAR_KEYWORD:
62                bound.arguments[name] = {key: wrap_node(arg) for key, arg in bound.arguments[name].items()}
63            else:
64                bound.arguments[name] = wrap_node(bound.arguments[name])
65        with TemporaryDirectory(prefix="dml-defunkify-") as tmpd:
66            with chdir(tmpd):
67                return fn(*bound.args, **bound.kwargs)
68
69    return wrapped