codekingpro/portable-devtools
114k
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 