Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion src/furu/__init__.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
from importlib.metadata import version

from furu.core import Furu
from furu.dependencies import dependency
from furu.dependencies import RecheckInterval, dependency
from furu.execution import load_or_create
from furu.logging import get_logger
from furu.migration import Migration
Expand All @@ -18,6 +18,7 @@
"LazyResult",
"Migration",
"ResourceRequirements",
"RecheckInterval",
"dependency",
"ResultCodec",
"ResultRegistry",
Expand Down
144 changes: 134 additions & 10 deletions src/furu/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,43 @@ class DagNode:
obj: Furu
dependencies: list[DagNode] = field(default_factory=list)
dependents: list[DagNode] = field(default_factory=list)
declared_dependency_ids: set[str] = field(default_factory=set)
runtime_dependency_ids: set[str] = field(default_factory=set)


def _ensure_edge(node: DagNode, dep_node: DagNode) -> None:
if dep_node not in node.dependencies:
node.dependencies.append(dep_node)
if node not in dep_node.dependents:
dep_node.dependents.append(node)


def _remove_edge(node: DagNode, dep_node: DagNode) -> None:
if dep_node in node.dependencies:
node.dependencies.remove(dep_node)
if node in dep_node.dependents:
dep_node.dependents.remove(node)


def _is_running(manager: Manager, node: DagNode) -> bool:
return any(running.node is node for running in manager.running.values())


def _move_ready_or_blocked(manager: Manager, node: DagNode) -> None:
object_id = node.obj.object_id
if (
object_id in manager.completed
or object_id in manager.failed
or _is_running(manager, node)
):
return

if node.dependencies:
manager.ready.pop(object_id, None)
manager.blocked[object_id] = node
else:
manager.blocked.pop(object_id, None)
manager.ready[object_id] = node


def _add_to_dag(manager: Manager, objs: Sequence[Furu]) -> None:
Expand Down Expand Up @@ -54,9 +91,13 @@ def _add_to_dag(manager: Manager, objs: Sequence[Furu]) -> None:
for obj_id, refs in refs_by_id.items():
node = manager.nodes_by_id[obj_id]
for ref in refs:
if ref.object_id in manager.completed:
continue
if dep_node := manager.nodes_by_id.get(ref.object_id):
node.dependencies.append(dep_node)
dep_node.dependents.append(node)
if dep_node.obj.status() == "completed":
continue
node.declared_dependency_ids.add(ref.object_id)
_ensure_edge(node, dep_node)

for node in newly_added:
if node.dependencies:
Expand Down Expand Up @@ -94,12 +135,95 @@ def _update_dag_blocking_dependencies(

for dependency_id in dependency_ids:
dep_node = manager.nodes_by_id[dependency_id]
if dep_node not in node.dependencies:
node.dependencies.append(dep_node)
if node not in dep_node.dependents:
dep_node.dependents.append(node)
node.runtime_dependency_ids.add(dependency_id)
_ensure_edge(node, dep_node)

if node.dependencies:
manager.blocked[node.obj.object_id] = node
else:
manager.ready[node.obj.object_id] = node
_move_ready_or_blocked(manager, node)


def _refresh_dag_declared_dependencies(manager: Manager) -> None:
nodes = tuple(manager.blocked.values()) + tuple(manager.ready.values())
for node in nodes:
if manager.nodes_by_id.get(node.obj.object_id) is not node:
continue
refs = collect_declared_refs(node.obj)
_reconcile_declared_dependencies(manager, node, refs)


def _reconcile_declared_dependencies(
manager: Manager,
node: DagNode,
refs: Sequence[Furu],
) -> None:
wanted_declared_ids: set[str] = set()
wanted_declared_ids_in_order: list[str] = []
missing_dependencies: list[Furu] = []

for ref in refs:
object_id = ref.object_id
if object_id == node.obj.object_id or object_id in manager.completed:
continue

dep_node = manager.nodes_by_id.get(object_id)
if dep_node is not None:
if dep_node.obj.status() != "completed":
wanted_declared_ids.add(object_id)
wanted_declared_ids_in_order.append(object_id)
continue

match ref.status():
case "completed":
continue
case "running" | "missing" | "failed":
wanted_declared_ids.add(object_id)
wanted_declared_ids_in_order.append(object_id)
missing_dependencies.append(ref)
case x:
assert_never(x)

_add_to_dag(manager, missing_dependencies)

for dependency_id in tuple(node.declared_dependency_ids - wanted_declared_ids):
node.declared_dependency_ids.discard(dependency_id)
if dependency_id in node.runtime_dependency_ids:
continue

dep_node = manager.nodes_by_id.get(dependency_id)
if dep_node is None:
continue

_remove_edge(node, dep_node)
_prune_orphan_dependency(manager, dep_node)

for dependency_id in wanted_declared_ids_in_order:
dep_node = manager.nodes_by_id.get(dependency_id)
if dep_node is None:
continue

node.declared_dependency_ids.add(dependency_id)
_ensure_edge(node, dep_node)

_move_ready_or_blocked(manager, node)


def _prune_orphan_dependency(manager: Manager, node: DagNode) -> None:
object_id = node.obj.object_id
if (
object_id in manager.root_ids
or node.dependents
or _is_running(manager, node)
or manager.nodes_by_id.get(object_id) is not node
):
return

manager.ready.pop(object_id, None)
manager.blocked.pop(object_id, None)
manager.completed.pop(object_id, None)
manager.failed.pop(object_id, None)
manager.nodes_by_id.pop(object_id, None)

for dep_node in tuple(node.dependencies):
node.declared_dependency_ids.discard(dep_node.obj.object_id)
node.runtime_dependency_ids.discard(dep_node.obj.object_id)
_remove_edge(node, dep_node)
_prune_orphan_dependency(manager, dep_node)
123 changes: 90 additions & 33 deletions src/furu/dependencies.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import fields, is_dataclass
from functools import cached_property
import time
from typing import TYPE_CHECKING, Any, Callable, Literal, overload

from pydantic import BaseModel as PydanticBaseModel
Expand All @@ -13,53 +13,110 @@
from furu.core import Furu


class _CachedDependency[T](cached_property):
__furu_dependency__ = True
type RecheckInterval = int | Literal["never"]
type _DependencyCacheEntry[T] = tuple[float | None, T]


class _UncachedDependency[T](property):
class _Dependency[T](property):
__furu_dependency__ = True


@overload
def dependency[TFuru: Furu[Any], T](
func: Callable[[TFuru], T], /
) -> _CachedDependency[T]: ...


@overload
def dependency[TFuru: Furu[Any], T](
*, cached: Literal[True] = True
) -> Callable[[Callable[[TFuru], T]], _CachedDependency[T]]: ...
def __init__(
self,
func: Callable[[Any], T],
*,
recheck_interval: RecheckInterval,
) -> None:
if isinstance(recheck_interval, bool):
raise TypeError(
"recheck_interval must be a non-negative integer or 'never'"
)
if isinstance(recheck_interval, int):
if recheck_interval < 0:
raise ValueError("recheck_interval must be non-negative")
elif recheck_interval != "never":
raise TypeError(
"recheck_interval must be a non-negative integer or 'never'"
)

super().__init__(func)
self._recheck_interval = recheck_interval
func_name = getattr(func, "__name__", type(func).__name__)
self._cache_key = f"__furu_dependency_cache_{func_name}"

def __set_name__(self, owner: type[object], name: str) -> None:
self._cache_key = f"__furu_dependency_cache_{name}"

@overload
def __get__(
self,
obj: None,
objtype: type[object] | None = None,
) -> _Dependency[T]: ...

@overload
def __get__(
self,
obj: object,
objtype: type[object] | None = None,
) -> T: ...

def __get__(
self,
obj: object | None,
objtype: type[object] | None = None,
) -> T | _Dependency[T]:
if obj is None:
return self

cached: _DependencyCacheEntry[T] | None = getattr(obj, self._cache_key, None)
now: float | None = None
if cached is not None:
expires_at, value = cached
if expires_at is None:
return value

now = time.monotonic()
if now < expires_at:
return value

if self.fget is None:
raise AttributeError("unreadable dependency")

recheck_interval = self._recheck_interval
if isinstance(recheck_interval, int) and now is None:
now = time.monotonic()

value = self.fget(obj)
if recheck_interval == "never":
expires_at = None
else:
assert isinstance(recheck_interval, int)
assert now is not None
expires_at = now + recheck_interval
object.__setattr__(obj, self._cache_key, (expires_at, value))
return value


@overload
def dependency[TFuru: Furu[Any], T](
*, cached: Literal[False]
) -> Callable[[Callable[[TFuru], T]], _UncachedDependency[T]]: ...
func: Callable[[TFuru], T], /
) -> _Dependency[T]: ...


@overload
def dependency[TFuru: Furu[Any], T](
*, cached: bool
) -> Callable[
[Callable[[TFuru], T]], _CachedDependency[T] | _UncachedDependency[T]
]: ...
*, recheck_interval: RecheckInterval = "never"
) -> Callable[[Callable[[TFuru], T]], _Dependency[T]]: ...


def dependency[TFuru: Furu[Any], T](
func: Callable[[TFuru], T] | None = None, /, *, cached: bool = True
) -> (
_CachedDependency[T]
| _UncachedDependency[T]
| Callable[[Callable[[TFuru], T]], _CachedDependency[T] | _UncachedDependency[T]]
):
def decorate(
func: Callable[[TFuru], T],
) -> _CachedDependency[T] | _UncachedDependency[T]:
if cached:
return _CachedDependency(func)
return _UncachedDependency(func)
func: Callable[[TFuru], T] | None = None,
/,
*,
recheck_interval: RecheckInterval = "never",
) -> _Dependency[T] | Callable[[Callable[[TFuru], T]], _Dependency[T]]:
def decorate(func: Callable[[TFuru], T]) -> _Dependency[T]:
return _Dependency(func, recheck_interval=recheck_interval)

if func is not None:
return decorate(func)
Expand Down
Loading
Loading