daggerml.contrib.api

View source
  1from __future__ import annotations
  2
  3import ast
  4import inspect
  5from dataclasses import dataclass, fields, is_dataclass, replace
  6from functools import wraps
  7from pathlib import Path
  8from textwrap import dedent
  9from typing import Any, Callable, Protocol, TypeAlias, TypeVar, cast, dataclass_transform, overload
 10
 11from daggerml import Runnable
 12from daggerml import api as core_api
 13from daggerml.api import DmlRepoError
 14from daggerml.contrib.codecs import DelayedLoad, DelayedRef, DelayedRunnable
 15
 16_DAGCLASS_CALL_NODE_NAME = "<dagclass-call>"
 17_DAGCLASS_RESERVED_NAMES = {f.name for f in fields(core_api.Dag)} | {
 18    name for name in dir(core_api.Dag) if not name.startswith("_")
 19}
 20
 21
 22def _iter_dagclass_members(instance):
 23    members = getattr(instance, "__dagclass_members__", None)
 24    order = getattr(instance, "__dagclass_member_order__", None)
 25    if not isinstance(members, dict) or not isinstance(order, list):
 26        raise DmlRepoError("dagclass instance is not compiled")
 27    for name in order:
 28        yield name, members[name]
 29
 30
 31class _DagclassAnalyzer:
 32    def __init__(self, *, member_names: set[str]):
 33        self.member_names = member_names
 34        self.dependencies: list[str] = []
 35        self._dep_set: set[str] = set()
 36        self.assignments: set[str] = set()
 37
 38    def _add_dependency(self, name: str) -> None:
 39        if name not in self._dep_set:
 40            self._dep_set.add(name)
 41            self.dependencies.append(name)
 42
 43    def _scan(self, node: ast.AST) -> None:
 44        if isinstance(node, ast.Attribute) and isinstance(node.value, ast.Name) and node.value.id == "self":
 45            if isinstance(node.ctx, ast.Load):
 46                self._add_dependency(node.attr)
 47            elif isinstance(node.ctx, ast.Store):
 48                self.assignments.add(node.attr)
 49        for child in ast.iter_child_nodes(node):
 50            self._scan(child)
 51
 52    def analyze(self, fn: ast.FunctionDef) -> list[str]:
 53        for statement in fn.body:
 54            self._scan(statement)
 55        reserved_assignments = self.assignments & _DAGCLASS_RESERVED_NAMES
 56        if reserved_assignments:
 57            bad = ", ".join(sorted(reserved_assignments))
 58            raise DmlRepoError(f"Cannot assign to reserved dagclass names: {bad}")
 59        dependencies = [
 60            name for name in self.dependencies if name not in self.assignments and name not in _DAGCLASS_RESERVED_NAMES
 61        ]
 62        for name in dependencies:
 63            if name not in self.member_names:
 64                raise DmlRepoError(f"Unknown dagclass member reference: self.{name}")
 65        return dependencies
 66
 67
 68def _analyze_dagclass_method(*, cls, method_name: str, method, member_names: set[str]):
 69    try:
 70        source = dedent(inspect.getsource(method))
 71    except (OSError, TypeError) as e:
 72        raise DmlRepoError(f"Failed to inspect dagclass method source for {cls.__name__}.{method_name}: {e}") from e
 73    module = ast.parse(source)
 74    if len(module.body) != 1 or not isinstance(module.body[0], ast.FunctionDef):
 75        raise DmlRepoError(f"dagclass method source for {cls.__name__}.{method_name} must be a single function")
 76    fn = module.body[0]
 77    if not fn.args.args or fn.args.args[0].arg != "self":
 78        raise DmlRepoError(f"dagclass method {cls.__name__}.{method_name} must declare self as first parameter")
 79    analyzer = _DagclassAnalyzer(member_names=member_names)
 80    return analyzer.analyze(fn), fn.decorator_list
 81
 82
 83def _compile_plain_dagclass_method(*, cls, method_name: str, method, member_names: set[str]):
 84    dependencies, decorators = _analyze_dagclass_method(
 85        cls=cls,
 86        method_name=method_name,
 87        method=method,
 88        member_names=member_names,
 89    )
 90    if decorators:
 91        raise DmlRepoError(f"dagclass method {cls.__name__}.{method_name} has unsupported decorators")
 92    delayed = funkify(method, uri="script", adapter="local", prepop={name: ref(name) for name in dependencies})
 93    return delayed, dependencies
 94
 95
 96def _dagclass_decorated_method(value: Any) -> Callable[..., Any] | None:
 97    current = value
 98    while isinstance(current, DelayedRunnable):
 99        fn = current.kwargs.get("fn")
100        if callable(fn):
101            params = list(inspect.signature(fn).parameters.values())
102            return fn if params and params[0].name == "self" else None
103        current = current.sub
104    return None
105
106
107def _add_dagclass_prepop(value: DelayedRunnable, dependencies: list[str]) -> DelayedRunnable:
108    if isinstance(value.sub, DelayedRunnable):
109        return replace(value, sub=_add_dagclass_prepop(value.sub, dependencies))
110    kwargs = dict(value.kwargs)
111    kwargs["prepop"] = {**kwargs.get("prepop", {}), **{name: ref(name) for name in dependencies}}
112    return replace(value, kwargs=kwargs)
113
114
115def _collect_member_dependencies(value: Any, member_names: set[str]) -> set[str]:
116    deps: set[str] = set()
117
118    def visit(obj: Any) -> None:
119        if isinstance(obj, DelayedRef):
120            if obj.name not in member_names:
121                raise DmlRepoError(f"Unknown dagclass member reference: {obj.name}")
122            deps.add(obj.name)
123            return
124        if isinstance(obj, DelayedLoad):
125            return
126        if isinstance(obj, DelayedRunnable):
127            visit(obj.sub)
128            visit(obj.kwargs)
129            return
130        if isinstance(obj, Runnable):
131            visit(obj.sub)
132            visit(obj.kwargs)
133            return
134        if isinstance(obj, dict):
135            for key, value in obj.items():
136                visit(key)
137                visit(value)
138            return
139        if isinstance(obj, (list, tuple, set, frozenset)):
140            for item in obj:
141                visit(item)
142            return
143
144    visit(value)
145    return deps
146
147
148def _toposort_members(member_deps: dict[str, set[str]], order_hint: list[str]) -> list[str]:
149    ordered: list[str] = []
150    temp: set[str] = set()
151    done: set[str] = set()
152
153    def visit(name: str) -> None:
154        if name in done:
155            return
156        if name in temp:
157            raise DmlRepoError(f"dagclass member dependency cycle detected at: {name}")
158        temp.add(name)
159        for dep in sorted(
160            member_deps.get(name, set()),
161            key=lambda item: order_hint.index(item) if item in order_hint else len(order_hint),
162        ):
163            visit(dep)
164        temp.remove(name)
165        done.add(name)
166        ordered.append(name)
167
168    for name in order_hint:
169        visit(name)
170    if set(ordered) != set(order_hint):
171        raise DmlRepoError("dagclass member ordering is incomplete or inconsistent")
172    return ordered
173
174
175def _bind_dagclass_member(value, members: dict[str, Any]):
176    if isinstance(value, DelayedRef):
177        if value.name not in members:
178            raise DmlRepoError(f"Unknown dagclass member reference: {value.name}")
179        return members[value.name]
180    if isinstance(value, DelayedRunnable):
181        return replace(
182            value,
183            sub=_bind_dagclass_member(value.sub, members),
184            kwargs={key: _bind_dagclass_member(item, members) for key, item in value.kwargs.items()},
185        )
186    if isinstance(value, Runnable):
187        return replace(
188            value,
189            sub=_bind_dagclass_member(value.sub, members),
190            kwargs={key: _bind_dagclass_member(item, members) for key, item in value.kwargs.items()},
191        )
192    if isinstance(value, dict):
193        return {key: _bind_dagclass_member(item, members) for key, item in value.items()}
194    if isinstance(value, list):
195        return [_bind_dagclass_member(item, members) for item in value]
196    if isinstance(value, tuple):
197        return tuple(_bind_dagclass_member(item, members) for item in value)
198    if isinstance(value, set):
199        return {_bind_dagclass_member(item, members) for item in value}
200    if isinstance(value, frozenset):
201        return frozenset(_bind_dagclass_member(item, members) for item in value)
202    return value
203
204
205def _bind_dagclass_value(value):
206    if getattr(value.__class__, "__dagclass__", False):
207        entrypoint = getattr(value.__class__, "__dagclass_entrypoint__", "main")
208        if not hasattr(value, entrypoint):
209            raise DmlRepoError(f"Dagclass instance missing configured entrypoint: {entrypoint}")
210        return value.__dagclass_members__[entrypoint]
211    return value
212
213
214FunkifyInput: TypeAlias = Callable[..., Any] | Runnable | DelayedRunnable
215DagclassType = TypeVar("DagclassType", bound=type[Any])
216
217
218class _DagclassProtocol(Protocol):
219    __dagclass__: bool
220    __dagclass_entrypoint__: str
221    __dagclass_wrapped_init__: bool
222
223    def __init__(self, *args: Any, **kwargs: Any) -> None: ...
224
225
226def is_node_like(x: object) -> bool:
227    """Return True if x is a Node or any Delayed* type (DelayedRef, DelayedLoad, DelayedRunnable)."""
228    return isinstance(x, (core_api.Node, DelayedRef, DelayedLoad, DelayedRunnable))
229
230
231def ref(name: str) -> DelayedRef:
232    return DelayedRef(name)
233
234
235def load(dagname: str, nodename: str | None = None) -> DelayedLoad:
236    return DelayedLoad(dagname=dagname, nodename=nodename)
237
238
239def _compile_dagclass_instance(instance) -> None:
240    if getattr(instance, "__dagclass_compiled__", False):
241        return
242
243    attributes: dict[str, Any] = {}
244    attribute_order: list[str] = []
245    data_descriptor_names: set[str] = set()
246    method_defs: dict[str, tuple[Any, DelayedRunnable | None]] = {}
247    field_names = {f.name for f in fields(instance)}
248
249    for f in fields(instance):
250        current = getattr(instance, f.name)
251        bound = _bind_dagclass_value(current)
252        if bound is not current:
253            setattr(instance, f.name, bound)
254        attributes[f.name] = getattr(instance, f.name)
255        attribute_order.append(f.name)
256
257    for name, class_value in instance.__class__.__dict__.items():
258        if name.startswith("_"):
259            continue
260        if name in field_names:
261            continue
262        if inspect.isdatadescriptor(cast(Any, class_value)):
263            data_descriptor_names.add(name)
264        if inspect.isfunction(class_value):
265            method_defs[name] = (class_value, None)
266            continue
267        decorated_method = _dagclass_decorated_method(class_value)
268        if decorated_method is not None:
269            method_defs[name] = (decorated_method, class_value)
270            continue
271        if name in instance.__dict__:
272            attributes[name] = getattr(instance, name)
273            attribute_order.append(name)
274            continue
275        bound = _bind_dagclass_value(class_value)
276        if bound is not class_value:
277            setattr(instance, name, bound)
278        attributes[name] = getattr(instance, name)
279        attribute_order.append(name)
280
281    member_names = set(attributes.keys()) | set(method_defs.keys())
282    method_names = set(method_defs.keys())
283    reserved = sorted(member_names & _DAGCLASS_RESERVED_NAMES)
284    if reserved:
285        bad = ", ".join(reserved)
286        raise DmlRepoError(f"dagclass uses reserved names: {bad}")
287    attribute_deps: dict[str, set[str]] = {}
288    attribute_names = set(attributes)
289    for name, value in attributes.items():
290        attribute_deps[name] = _collect_member_dependencies(value, attribute_names)
291
292    members: dict[str, Any] = {}
293    order: list[str] = []
294    for name in _toposort_members(attribute_deps, attribute_order):
295        bound = _bind_dagclass_member(attributes[name], members)
296        if name not in data_descriptor_names:
297            setattr(instance, name, bound)
298        members[name] = bound
299        order.append(name)
300
301    compiled_methods: dict[str, Any] = {}
302    method_deps: dict[str, set[str]] = {}
303    for name, (method, decorated) in method_defs.items():
304        if decorated is None:
305            compiled, deps = _compile_plain_dagclass_method(
306                cls=instance.__class__,
307                method_name=name,
308                method=method,
309                member_names=member_names,
310            )
311        else:
312            deps, _decorators = _analyze_dagclass_method(
313                cls=instance.__class__,
314                method_name=name,
315                method=method,
316                member_names=member_names,
317            )
318            compiled = _add_dagclass_prepop(decorated, deps)
319        compiled_methods[name] = compiled
320        method_deps[name] = set(deps) & method_names
321
322    method_order = _toposort_members(method_deps, list(method_defs))
323    for name in method_order:
324        compiled = _bind_dagclass_member(compiled_methods[name], members)
325        setattr(instance, name, compiled)
326        members[name] = compiled
327        order.append(name)
328
329    instance.__dagclass_members__ = members
330    instance.__dagclass_member_order__ = order
331    instance.__dagclass_compiled__ = True
332
333
334@dataclass_transform()
335@overload
336def dagclass(
337    _cls: None = None, *, entrypoint: str = "main", **dataclass_kwargs: Any
338) -> Callable[[DagclassType], DagclassType]: ...
339@overload
340def dagclass(_cls: DagclassType, *, entrypoint: str = "main", **dataclass_kwargs: Any) -> DagclassType: ...
341def dagclass(
342    _cls: DagclassType | None = None, *, entrypoint: str = "main", **dataclass_kwargs: Any
343) -> Callable[[DagclassType], DagclassType] | DagclassType:
344    def wrap(cls: DagclassType) -> DagclassType:
345        if not is_dataclass(cls):
346            cls = dataclass(cls, **dataclass_kwargs)
347        elif dataclass_kwargs:
348            bad = ", ".join(sorted(dataclass_kwargs.keys()))
349            raise DmlRepoError(f"api.dagclass dataclass kwargs not allowed on pre-dataclass class: {bad}")
350        cls = cast(DagclassType, cls)
351        dagclass_cls = cast(_DagclassProtocol, cls)
352        dagclass_cls.__dagclass__ = True
353        dagclass_cls.__dagclass_entrypoint__ = entrypoint
354        if getattr(dagclass_cls, "__dagclass_wrapped_init__", False):
355            return cls
356        original_init = dagclass_cls.__init__
357
358        @wraps(original_init)
359        def _dagclass_init(self, *args, **kwargs):
360            original_init(self, *args, **kwargs)
361            _compile_dagclass_instance(self)
362
363        dagclass_cls.__init__ = _dagclass_init
364        dagclass_cls.__dagclass_wrapped_init__ = True
365        return cls
366
367    if _cls is None:
368        return wrap
369    return wrap(_cls)
370
371
372def _default_run_name(instance) -> str:
373    module = __import__(instance.__class__.__module__, fromlist=["__name__"])
374    if not getattr(module, "__file__", None):
375        return f"{instance.__class__.__module__}::{instance.__class__.__name__}"
376    module_file = Path(module.__file__).resolve()
377    repo_root = None
378    for parent in (module_file.parent, *module_file.parents):
379        if (parent / ".git").exists():
380            repo_root = parent
381            break
382    base = repo_root if repo_root is not None else Path.cwd().resolve()
383    try:
384        rel = module_file.relative_to(base)
385    except ValueError:
386        rel = module_file
387    rel_no_ext = rel.with_suffix("").as_posix()
388    return f"{rel_no_ext}::{instance.__class__.__name__}"
389
390
391def run(instance, *args, name: str | None = None, entrypoint: str | None = None, **kwargs):
392    if not getattr(instance.__class__, "__dagclass__", False):
393        raise DmlRepoError("api.run instance is not a dagclass instance")
394    if not getattr(instance, "__dagclass_compiled__", False):
395        raise DmlRepoError("api.run instance is not compiled")
396    entry = entrypoint or getattr(instance.__class__, "__dagclass_entrypoint__", "main")
397    if not hasattr(instance, entry):
398        raise DmlRepoError(f"api.run entrypoint not found: {entry}")
399    fn = instance.__dagclass_members__.get(entry)
400    if not isinstance(fn, DelayedRunnable):
401        raise DmlRepoError("api.run entrypoint must be DelayedRunnable")
402    run_name = name or _default_run_name(instance)
403    dml = core_api.get_default_dml()
404    with core_api.new(dml=dml, name=run_name, message=run_name) as dag:
405        result = dag.call(fn, *args, name=_DAGCLASS_CALL_NODE_NAME, **kwargs)
406        dag.commit(result)
407
408
409@overload
410def funkify(
411    sub_or_fn: None = None, *, adapter: str = "local", uri: str = "script", **kwargs: Any
412) -> Callable[[FunkifyInput], DelayedRunnable]: ...
413@overload
414def funkify(
415    sub_or_fn: Callable[..., Any], *, adapter: str = "local", uri: str = "script", **kwargs: Any
416) -> DelayedRunnable: ...
417@overload
418def funkify(
419    sub_or_fn: Runnable | DelayedRunnable, *, adapter: str = "local", uri: str = "script", **kwargs: Any
420) -> DelayedRunnable: ...
421def funkify(
422    sub_or_fn: FunkifyInput | None = None, *, adapter: str = "local", uri: str = "script", **kwargs: Any
423) -> Callable[[FunkifyInput], DelayedRunnable] | DelayedRunnable:
424    def _make(value: FunkifyInput) -> DelayedRunnable:
425        if callable(value):
426            reserved = sorted({"fn", "script", "fn_name"} & kwargs.keys())
427            if reserved:
428                raise DmlRepoError(f"Unknown kwarg: {reserved[0]}")
429            delayed_kwargs = {"fn": value, **kwargs}
430            if uri == "script":
431                from daggerml.contrib.executors.script import ScriptExecutor
432
433                script = ScriptExecutor._render_script(
434                    value,
435                    extra_objs=list(kwargs.get("extra_objs", [])),
436                    post_lines=list(kwargs.get("post_lines", [])),
437                )
438                delayed_kwargs["script"] = script
439                delayed_kwargs["fn_name"] = value.__name__
440            return DelayedRunnable(uri=uri, adapter=adapter, sub=None, kwargs=delayed_kwargs)
441        if isinstance(value, (Runnable, DelayedRunnable)):
442            return DelayedRunnable(uri=uri, adapter=adapter, sub=value, kwargs=dict(kwargs))
443        raise DmlRepoError(f"Invalid funkify input: {type(value).__name__}")
444
445    if sub_or_fn is None:
446        return _make
447    return _make(sub_or_fn)

FunkifyInput

FunkifyInput: TypeAlias= Union[Callable[..., Any], daggerml.Runnable, daggerml.contrib.codecs.DelayedRunnable]

is_node_like

def is_node_like(x: object) -> bool:
View source
227def is_node_like(x: object) -> bool:
228    """Return True if x is a Node or any Delayed* type (DelayedRef, DelayedLoad, DelayedRunnable)."""
229    return isinstance(x, (core_api.Node, DelayedRef, DelayedLoad, DelayedRunnable))

Return True if x is a Node or any Delayed* type (DelayedRef, DelayedLoad, DelayedRunnable).

ref

def ref(name: str) -> daggerml.contrib.codecs.DelayedRef:
View source
232def ref(name: str) -> DelayedRef:
233    return DelayedRef(name)

load

def load( dagname: str, nodename: str | None = None) -> daggerml.contrib.codecs.DelayedLoad:
View source
236def load(dagname: str, nodename: str | None = None) -> DelayedLoad:
237    return DelayedLoad(dagname=dagname, nodename=nodename)

dagclass

def dagclass( _cls: Optional[~DagclassType] = None, *, entrypoint: str = 'main', **dataclass_kwargs: Any) -> Union[Callable[[~DagclassType], ~DagclassType], ~DagclassType]:
View source
342def dagclass(
343    _cls: DagclassType | None = None, *, entrypoint: str = "main", **dataclass_kwargs: Any
344) -> Callable[[DagclassType], DagclassType] | DagclassType:
345    def wrap(cls: DagclassType) -> DagclassType:
346        if not is_dataclass(cls):
347            cls = dataclass(cls, **dataclass_kwargs)
348        elif dataclass_kwargs:
349            bad = ", ".join(sorted(dataclass_kwargs.keys()))
350            raise DmlRepoError(f"api.dagclass dataclass kwargs not allowed on pre-dataclass class: {bad}")
351        cls = cast(DagclassType, cls)
352        dagclass_cls = cast(_DagclassProtocol, cls)
353        dagclass_cls.__dagclass__ = True
354        dagclass_cls.__dagclass_entrypoint__ = entrypoint
355        if getattr(dagclass_cls, "__dagclass_wrapped_init__", False):
356            return cls
357        original_init = dagclass_cls.__init__
358
359        @wraps(original_init)
360        def _dagclass_init(self, *args, **kwargs):
361            original_init(self, *args, **kwargs)
362            _compile_dagclass_instance(self)
363
364        dagclass_cls.__init__ = _dagclass_init
365        dagclass_cls.__dagclass_wrapped_init__ = True
366        return cls
367
368    if _cls is None:
369        return wrap
370    return wrap(_cls)

run

def run( instance, *args, name: str | None = None, entrypoint: str | None = None, **kwargs):
View source
392def run(instance, *args, name: str | None = None, entrypoint: str | None = None, **kwargs):
393    if not getattr(instance.__class__, "__dagclass__", False):
394        raise DmlRepoError("api.run instance is not a dagclass instance")
395    if not getattr(instance, "__dagclass_compiled__", False):
396        raise DmlRepoError("api.run instance is not compiled")
397    entry = entrypoint or getattr(instance.__class__, "__dagclass_entrypoint__", "main")
398    if not hasattr(instance, entry):
399        raise DmlRepoError(f"api.run entrypoint not found: {entry}")
400    fn = instance.__dagclass_members__.get(entry)
401    if not isinstance(fn, DelayedRunnable):
402        raise DmlRepoError("api.run entrypoint must be DelayedRunnable")
403    run_name = name or _default_run_name(instance)
404    dml = core_api.get_default_dml()
405    with core_api.new(dml=dml, name=run_name, message=run_name) as dag:
406        result = dag.call(fn, *args, name=_DAGCLASS_CALL_NODE_NAME, **kwargs)
407        dag.commit(result)

funkify

def funkify( sub_or_fn: Union[Callable[..., Any], daggerml.Runnable, daggerml.contrib.codecs.DelayedRunnable, NoneType] = None, *, adapter: str = 'local', uri: str = 'script', **kwargs: Any) -> Union[Callable[[Union[Callable[..., Any], daggerml.Runnable, daggerml.contrib.codecs.DelayedRunnable]], daggerml.contrib.codecs.DelayedRunnable], daggerml.contrib.codecs.DelayedRunnable]:
View source
422def funkify(
423    sub_or_fn: FunkifyInput | None = None, *, adapter: str = "local", uri: str = "script", **kwargs: Any
424) -> Callable[[FunkifyInput], DelayedRunnable] | DelayedRunnable:
425    def _make(value: FunkifyInput) -> DelayedRunnable:
426        if callable(value):
427            reserved = sorted({"fn", "script", "fn_name"} & kwargs.keys())
428            if reserved:
429                raise DmlRepoError(f"Unknown kwarg: {reserved[0]}")
430            delayed_kwargs = {"fn": value, **kwargs}
431            if uri == "script":
432                from daggerml.contrib.executors.script import ScriptExecutor
433
434                script = ScriptExecutor._render_script(
435                    value,
436                    extra_objs=list(kwargs.get("extra_objs", [])),
437                    post_lines=list(kwargs.get("post_lines", [])),
438                )
439                delayed_kwargs["script"] = script
440                delayed_kwargs["fn_name"] = value.__name__
441            return DelayedRunnable(uri=uri, adapter=adapter, sub=None, kwargs=delayed_kwargs)
442        if isinstance(value, (Runnable, DelayedRunnable)):
443            return DelayedRunnable(uri=uri, adapter=adapter, sub=value, kwargs=dict(kwargs))
444        raise DmlRepoError(f"Invalid funkify input: {type(value).__name__}")
445
446    if sub_or_fn is None:
447        return _make
448    return _make(sub_or_fn)