|
1 | 1 | from __future__ import annotations |
2 | | -from typing import Any, Protocol, Type |
3 | | -from graphlib import TopologicalSorter, CycleError |
| 2 | +from typing import Any, Protocol |
4 | 3 | from graphai.callback import Callback |
5 | 4 | from graphai.utils import logger |
6 | 5 |
|
@@ -63,7 +62,7 @@ def __init__( |
63 | 62 | self.edges: list[Any] = [] |
64 | 63 | self.start_node: NodeProtocol | None = None |
65 | 64 | self.end_nodes: list[NodeProtocol] = [] |
66 | | - self.Callback: Type[Callback] = Callback |
| 65 | + self.Callback: type[Callback] = Callback |
67 | 66 | self.max_steps = max_steps |
68 | 67 | self.state = initial_state or {} |
69 | 68 |
|
@@ -288,17 +287,6 @@ def _add_edge(src: str, dst: str) -> None: |
288 | 287 | "(src, Iterable[dst]), mapping{'source'/'destination'}, or objects with .source/.destination" |
289 | 288 | ) |
290 | 289 |
|
291 | | - # cycle detection |
292 | | - preds: dict[str, set[str]] = {n: set() for n in nodes.keys()} |
293 | | - for s, ds in adj.items(): |
294 | | - for d in ds: |
295 | | - preds[d].add(s) |
296 | | - |
297 | | - try: |
298 | | - list(TopologicalSorter(preds).static_order()) |
299 | | - except CycleError as e: |
300 | | - raise GraphCompileError("Cycle detected in graph") from e |
301 | | - |
302 | 290 | # reachability from start |
303 | 291 | seen: set[str] = set() |
304 | 292 | stack = [start_name] |
@@ -388,7 +376,7 @@ def set_callback(self, callback_class: type[Callback]) -> "Graph": |
388 | 376 | as the default callback when no callback is passed to the `execute` method. |
389 | 377 |
|
390 | 378 | :param callback_class: The callback class to use as the default callback. |
391 | | - :type callback_class: Type[Callback] |
| 379 | + :type callback_class: type[Callback] |
392 | 380 | """ |
393 | 381 | self.Callback = callback_class |
394 | 382 | return self |
|
0 commit comments