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)
Return True if x is a Node or any Delayed* type (DelayedRef, DelayedLoad, DelayedRunnable).
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)
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)
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)