Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_utils.py292 linesDownload Raw Back to pregel
1from __future__ import annotations2 3import ast4import inspect5import re6import textwrap7from collections.abc import Callable, Sequence8from functools import partial9from typing import Any10 11from langchain_core.runnables import (12    Runnable,13    RunnableLambda,14    RunnableParallel,15    RunnableSequence,16)17from langchain_core.runnables.base import RunnableBindingBase18from langchain_core.runnables.config import run_in_executor19from langgraph.checkpoint.base import ChannelVersions20from typing_extensions import override21 22from langgraph._internal._runnable import RunnableCallable, RunnableSeq23from langgraph._internal._timeout import sync_timeout_unsupported24from langgraph.pregel.protocol import PregelProtocol25 26_SEQUENCE_TYPES = (RunnableSeq, RunnableSequence)27 28 29def get_new_channel_versions(30    previous_versions: ChannelVersions, current_versions: ChannelVersions31) -> ChannelVersions:32    """Get subset of current_versions that are newer than previous_versions."""33    if previous_versions:34        version_type = type(next(iter(current_versions.values()), None))35        null_version = version_type()  # type: ignore[misc]36        new_versions = {37            k: v38            for k, v in current_versions.items()39            if v > previous_versions.get(k, null_version)  # type: ignore[operator]40        }41    else:42        new_versions = current_versions43 44    return new_versions45 46 47def find_subgraph_pregel(candidate: Runnable) -> PregelProtocol | None:48    from langgraph.pregel import Pregel49 50    candidates: list[Runnable] = [candidate]51 52    for c in candidates:53        if (54            isinstance(c, PregelProtocol)55            # subgraphs that disabled checkpointing are not considered56            and (not isinstance(c, Pregel) or c.checkpointer is not False)57        ):58            return c59        elif isinstance(c, RunnableSequence) or isinstance(c, RunnableSeq):60            candidates.extend(c.steps)61        elif isinstance(c, RunnableLambda):62            candidates.extend(c.deps)63        elif isinstance(c, RunnableCallable):64            if c.func is not None:65                candidates.extend(66                    nl.__self__ if hasattr(nl, "__self__") else nl67                    for nl in get_function_nonlocals(c.func)68                )69            elif c.afunc is not None:70                candidates.extend(71                    nl.__self__ if hasattr(nl, "__self__") else nl72                    for nl in get_function_nonlocals(c.afunc)73                )74 75    return None76 77 78def _sequence_steps(runnable: Runnable) -> Sequence[Runnable] | None:79    if isinstance(runnable, _SEQUENCE_TYPES):80        return runnable.steps81    return None82 83 84def _parallel_steps(runnable: Runnable) -> Sequence[Runnable] | None:85    if isinstance(runnable, RunnableParallel):86        return tuple(runnable.steps__.values())87    return None88 89 90def _has_method_override(runnable: Runnable, method_name: str) -> bool:91    method = getattr(type(runnable), method_name, None)92    return method is not None and method is not getattr(Runnable, method_name)93 94 95def _is_executor_backed_afunc(afunc: Callable[..., Any] | None) -> bool:96    return isinstance(afunc, partial) and afunc.func is run_in_executor97 98 99def _has_native_async(runnable: Runnable) -> bool:100    if isinstance(runnable, RunnableCallable):101        return runnable.afunc is not None and not _is_executor_backed_afunc(102            runnable.afunc103        )104    if isinstance(runnable, RunnableLambda):105        return bool(getattr(runnable, "afunc", False))106    return _has_method_override(runnable, "ainvoke")107 108 109def _runnable_has_native_async(runnable: Runnable) -> bool:110    """Return whether a runnable can be idle-timed without known sync code.111 112    For custom runnable subclasses, an `ainvoke` override is treated as the113    async contract. We do not introspect whether that implementation delegates114    to blocking work internally — e.g. a subclass whose `ainvoke` calls115    `asyncio.to_thread(self.invoke, ...)` will pass this check but the wrapped116    sync work is still uncancellable. Idle-timeout enforcement on such a117    runnable will fire `NodeTimeoutError` correctly, but the background thread118    will keep running until its sync work returns.119    """120 121    while isinstance(runnable, RunnableBindingBase):122        runnable = runnable.bound123    steps = _sequence_steps(runnable)124    if steps is None:125        steps = _parallel_steps(runnable)126    if steps is not None:127        return all(_runnable_has_native_async(step) for step in steps)128    # Raw callables and the common composition wrappers created by graph129    # builders fall through here. We do not exhaustively unwrap every Runnable130    # wrapper — wrappers that provide `ainvoke` are treated as owning the async131    # contract.132    return _has_native_async(runnable)133 134 135def validate_timeout_supported(runnable: Runnable, *, name: str) -> None:136    if not _runnable_has_native_async(runnable):137        raise sync_timeout_unsupported(name)138 139 140def get_function_nonlocals(func: Callable) -> list[Any]:141    """Get the nonlocal variables accessed by a function.142 143    Args:144        func: The function to check.145 146    Returns:147        List[Any]: The nonlocal variables accessed by the function.148    """149    try:150        code = inspect.getsource(func)151        tree = ast.parse(textwrap.dedent(code))152        visitor = FunctionNonLocals()153        visitor.visit(tree)154        values: list[Any] = []155        closure = (156            inspect.getclosurevars(func.__wrapped__)157            if hasattr(func, "__wrapped__") and callable(func.__wrapped__)158            else inspect.getclosurevars(func)159        )160        candidates = {**closure.globals, **closure.nonlocals}161        for k, v in candidates.items():162            if k in visitor.nonlocals:163                values.append(v)164            for kk in visitor.nonlocals:165                if "." in kk and kk.startswith(k):166                    vv = v167                    for part in kk.split(".")[1:]:168                        if vv is None:169                            break170                        else:171                            try:172                                vv = getattr(vv, part)173                            except AttributeError:174                                break175                    else:176                        values.append(vv)177    except (SyntaxError, TypeError, OSError, SystemError):178        return []179 180    return values181 182 183class FunctionNonLocals(ast.NodeVisitor):184    """Get the nonlocal variables accessed of a function."""185 186    def __init__(self) -> None:187        self.nonlocals: set[str] = set()188 189    @override190    def visit_FunctionDef(self, node: ast.FunctionDef) -> Any:191        """Visit a function definition.192 193        Args:194            node: The node to visit.195 196        Returns:197            Any: The result of the visit.198        """199        visitor = NonLocals()200        visitor.visit(node)201        self.nonlocals.update(visitor.loads - visitor.stores)202 203    @override204    def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> Any:205        """Visit an async function definition.206 207        Args:208            node: The node to visit.209 210        Returns:211            Any: The result of the visit.212        """213        visitor = NonLocals()214        visitor.visit(node)215        self.nonlocals.update(visitor.loads - visitor.stores)216 217    @override218    def visit_Lambda(self, node: ast.Lambda) -> Any:219        """Visit a lambda function.220 221        Args:222            node: The node to visit.223 224        Returns:225            Any: The result of the visit.226        """227        visitor = NonLocals()228        visitor.visit(node)229        self.nonlocals.update(visitor.loads - visitor.stores)230 231 232class NonLocals(ast.NodeVisitor):233    """Get nonlocal variables accessed."""234 235    def __init__(self) -> None:236        self.loads: set[str] = set()237        self.stores: set[str] = set()238 239    @override240    def visit_Name(self, node: ast.Name) -> Any:241        """Visit a name node.242 243        Args:244            node: The node to visit.245 246        Returns:247            Any: The result of the visit.248        """249        if isinstance(node.ctx, ast.Load):250            self.loads.add(node.id)251        elif isinstance(node.ctx, ast.Store):252            self.stores.add(node.id)253 254    @override255    def visit_Attribute(self, node: ast.Attribute) -> Any:256        """Visit an attribute node.257 258        Args:259            node: The node to visit.260 261        Returns:262            Any: The result of the visit.263        """264        if isinstance(node.ctx, ast.Load):265            parent = node.value266            attr_expr = node.attr267            while isinstance(parent, ast.Attribute):268                attr_expr = parent.attr + "." + attr_expr269                parent = parent.value270            if isinstance(parent, ast.Name):271                self.loads.add(parent.id + "." + attr_expr)272                self.loads.discard(parent.id)273            elif isinstance(parent, ast.Call):274                if isinstance(parent.func, ast.Name):275                    self.loads.add(parent.func.id)276                else:277                    parent = parent.func278                    attr_expr = ""279                    while isinstance(parent, ast.Attribute):280                        if attr_expr:281                            attr_expr = parent.attr + "." + attr_expr282                        else:283                            attr_expr = parent.attr284                        parent = parent.value285                    if isinstance(parent, ast.Name):286                        self.loads.add(parent.id + "." + attr_expr)287 288 289def is_xxh3_128_hexdigest(value: str) -> bool:290    """Check if the given string matches the format of xxh3_128_hexdigest."""291    return bool(re.fullmatch(r"[0-9a-f]{32}", value))292 
codekingpro/portable-devtools · Team Ai