Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
__init__.py621 linesDownload Raw Back to func
1from __future__ import annotations2 3import functools4import inspect5import warnings6from collections.abc import Awaitable, Callable, Sequence7from dataclasses import dataclass8from datetime import timedelta9from typing import (10    Any,11    Generic,12    TypeVar,13    cast,14    get_args,15    get_origin,16    overload,17)18 19from langgraph.cache.base import BaseCache20from langgraph.checkpoint.base import BaseCheckpointSaver21from langgraph.store.base import BaseStore22from typing_extensions import Unpack23 24from langgraph._internal import _serde25from langgraph._internal._constants import CACHE_NS_WRITES, PREVIOUS26from langgraph._internal._runnable import is_async_callable27from langgraph._internal._timeout import (28    coerce_timeout_policy,29    sync_timeout_unsupported,30)31from langgraph._internal._typing import MISSING, DeprecatedKwargs32from langgraph.channels.ephemeral_value import EphemeralValue33from langgraph.channels.last_value import LastValue34from langgraph.constants import END, START35from langgraph.pregel import Pregel36from langgraph.pregel._call import (37    P,38    SyncAsyncFuture,39    T,40    _call_with_options,41    get_runnable_for_entrypoint,42    identifier,43)44from langgraph.pregel._read import PregelNode45from langgraph.pregel._write import ChannelWrite, ChannelWriteEntry46from langgraph.types import (47    _DC_KWARGS,48    CachePolicy,49    RetryPolicy,50    StreamMode,51    TimeoutPolicy,52)53from langgraph.typing import ContextT54from langgraph.warnings import LangGraphDeprecatedSinceV05, LangGraphDeprecatedSinceV1055 56__all__ = ("task", "entrypoint")57 58 59class _TaskFunction(Generic[P, T]):60    def __init__(61        self,62        func: Callable[P, Awaitable[T]] | Callable[P, T],63        *,64        retry_policy: Sequence[RetryPolicy],65        cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None,66        timeout: TimeoutPolicy | None = None,67        name: str | None = None,68    ) -> None:69        if name is not None:70            if hasattr(func, "__func__"):71                # handle class methods72                # NOTE: we're modifying the instance method to avoid modifying73                # the original class method in case it's shared across multiple tasks74                instance_method = functools.partial(func.__func__, func.__self__)  # type: ignore [union-attr]75                instance_method.__name__ = name  # type: ignore [attr-defined]76                func = instance_method77            else:78                # handle regular functions / partials / callable classes, etc.79                func.__name__ = name80        self.func = func81        self.retry_policy = retry_policy82        self.cache_policy = cache_policy83        self.timeout = timeout84        functools.update_wrapper(self, func)85 86    def __call__(self, *args: P.args, **kwargs: P.kwargs) -> SyncAsyncFuture[T]:87        return _call_with_options(88            self.func,89            args,90            kwargs,91            retry_policy=self.retry_policy,92            cache_policy=self.cache_policy,93            timeout=self.timeout,94        )95 96    def clear_cache(self, cache: BaseCache) -> None:97        """Clear the cache for this task."""98        if self.cache_policy is not None:99            cache.clear(((CACHE_NS_WRITES, identifier(self.func) or "__dynamic__"),))100 101    async def aclear_cache(self, cache: BaseCache) -> None:102        """Clear the cache for this task."""103        if self.cache_policy is not None:104            await cache.aclear(105                ((CACHE_NS_WRITES, identifier(self.func) or "__dynamic__"),)106            )107 108 109@overload110def task(111    __func_or_none__: None = None,112    *,113    name: str | None = None,114    retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,115    cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None,116    timeout: float | timedelta | TimeoutPolicy | None = None,117    **kwargs: Unpack[DeprecatedKwargs],118) -> Callable[119    [Callable[P, Awaitable[T]] | Callable[P, T]],120    _TaskFunction[P, T],121]: ...122 123 124@overload125def task(__func_or_none__: Callable[P, Awaitable[T]]) -> _TaskFunction[P, T]: ...126 127 128@overload129def task(__func_or_none__: Callable[P, T]) -> _TaskFunction[P, T]: ...130 131 132def task(133    __func_or_none__: Callable[P, Awaitable[T]] | Callable[P, T] | None = None,134    *,135    name: str | None = None,136    retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,137    cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None,138    timeout: float | timedelta | TimeoutPolicy | None = None,139    **kwargs: Unpack[DeprecatedKwargs],140) -> (141    Callable[[Callable[P, Awaitable[T]] | Callable[P, T]], _TaskFunction[P, T]]142    | _TaskFunction[P, T]143):144    """Define a LangGraph task using the `task` decorator.145 146    !!! important "Requires python 3.11 or higher for async functions"147        The `task` decorator supports both sync and async functions. To use async148        functions, ensure that you are using Python 3.11 or higher.149 150    Tasks can only be called from within an [`entrypoint`][langgraph.func.entrypoint] or151    from within a `StateGraph`. A task can be called like a regular function with the152    following differences:153 154    - When a checkpointer is enabled, the function inputs and outputs must be serializable.155    - The decorated function can only be called from within an entrypoint or `StateGraph`.156    - Calling the function produces a future. This makes it easy to parallelize tasks.157 158    Args:159        name: An optional name for the task. If not provided, the function name will be used.160        retry_policy: An optional retry policy (or list of policies) to use for the task in case of a failure.161        cache_policy: An optional cache policy to use for the task. This allows caching of the task results.162        timeout: Timeout for each task attempt. A number or `timedelta` is a hard163            wall-clock cap and is not refreshed. Use `TimeoutPolicy` to configure164            both a wall-clock `run_timeout` and an `idle_timeout` refreshed by165            progress signals. For long-running work that doesn't naturally emit166            progress, call `runtime.heartbeat()` from inside the task. When the167            timeout fires, `NodeTimeoutError` is raised and the retry policy (if168            any) decides whether to retry. Supported only for async tasks; sync169            tasks cannot be safely cancelled in-process.170 171    Returns:172        A callable function when used as a decorator.173 174    Example: Sync Task175        ```python176        from langgraph.func import entrypoint, task177 178 179        @task180        def add_one_task(a: int) -> int:181            return a + 1182 183 184        @entrypoint()185        def add_one(numbers: list[int]) -> list[int]:186            futures = [add_one_task(n) for n in numbers]187            results = [f.result() for f in futures]188            return results189 190 191        # Call the entrypoint192        add_one.invoke([1, 2, 3])  # Returns [2, 3, 4]193        ```194 195    Example: Async Task196        ```python197        import asyncio198        from langgraph.func import entrypoint, task199 200 201        @task202        async def add_one_task(a: int) -> int:203            return a + 1204 205 206        @entrypoint()207        async def add_one(numbers: list[int]) -> list[int]:208            futures = [add_one_task(n) for n in numbers]209            return asyncio.gather(*futures)210 211 212        # Call the entrypoint213        await add_one.ainvoke([1, 2, 3])  # Returns [2, 3, 4]214        ```215    """216    if (retry := kwargs.get("retry", MISSING)) is not MISSING:217        warnings.warn(218            "`retry` is deprecated and will be removed. Please use `retry_policy` instead.",219            category=LangGraphDeprecatedSinceV05,220            stacklevel=2,221        )222        if retry_policy is None:223            retry_policy = retry  # type: ignore[assignment]224    timeout_policy = coerce_timeout_policy(timeout)225 226    retry_policies: Sequence[RetryPolicy] = (227        ()228        if retry_policy is None229        else (retry_policy,)230        if isinstance(retry_policy, RetryPolicy)231        else retry_policy232    )233 234    def decorator(235        func: Callable[P, Awaitable[T]] | Callable[P, T],236    ) -> Callable[P, SyncAsyncFuture[T]]:237        if timeout_policy is not None and not is_async_callable(func):238            name_ = name or getattr(func, "__name__", func.__class__.__name__)239            raise sync_timeout_unsupported(str(name_), kind="Task")240        return _TaskFunction(241            func,242            retry_policy=retry_policies,243            cache_policy=cache_policy,244            timeout=timeout_policy,245            name=name,246        )247 248    if __func_or_none__ is not None:249        return decorator(__func_or_none__)250 251    return decorator252 253 254R = TypeVar("R")255S = TypeVar("S")256 257 258# The decorator was wrapped in a class to support the `final` attribute.259# In this form, the `final` attribute should play nicely with IDE autocompletion,260# and type checking tools.261# In addition, we'll be able to surface this information in the API Reference.262class entrypoint(Generic[ContextT]):263    """Define a LangGraph workflow using the `entrypoint` decorator.264 265    ### Function signature266 267    The decorated function must accept a **single parameter**, which serves as the input268    to the function. This input parameter can be of any type. Use a dictionary269    to pass **multiple parameters** to the function.270 271    ### Injectable parameters272 273    The decorated function can request access to additional parameters274    that will be injected automatically at run time. These parameters include:275 276    | Parameter        | Description                                                                                          |277    |------------------|------------------------------------------------------------------------------------------------------|278    | **`config`**     | A configuration object (aka `RunnableConfig`) that holds run-time configuration values.              |279    | **`previous`**   | The previous return value for the given thread (available only when a checkpointer is provided).     |280    | **`runtime`**    | A `Runtime` object that contains information about the current run, including context, store, writer |281 282    The entrypoint decorator can be applied to sync functions or async functions.283 284    ### State management285 286    The **`previous`** parameter can be used to access the return value of the previous287    invocation of the entrypoint on the same thread id. This value is only available288    when a checkpointer is provided.289 290    If you want **`previous`** to be different from the return value, you can use the291    `entrypoint.final` object to return a value while saving a different value to the292    checkpoint.293 294    Args:295        checkpointer: Specify a checkpointer to create a workflow that can persist296            its state across runs.297        store: A generalized key-value store. Some implementations may support298            semantic search capabilities through an optional `index` configuration.299        cache: A cache to use for caching the results of the workflow.300        context_schema: Specifies the schema for the context object that will be301            passed to the workflow.302        cache_policy: A cache policy to use for caching the results of the workflow.303        retry_policy: A retry policy (or list of policies) to use for the workflow in case of a failure.304        timeout: Timeout for each workflow attempt. A number or `timedelta` is a305            hard wall-clock cap and is not refreshed. Use `TimeoutPolicy` to306            configure both a wall-clock `run_timeout` and an `idle_timeout`307            refreshed by progress signals. For long-running work that doesn't308            naturally emit progress, call `runtime.heartbeat()` from inside the309            workflow. When the timeout fires, `NodeTimeoutError` is raised and310            the retry policy (if any) decides whether to retry. Supported only311            for async workflows; sync workflows cannot be safely cancelled312            in-process.313 314    !!! warning "`config_schema` Deprecated"315        The `config_schema` parameter is deprecated in v0.6.0 and support will be removed in v2.0.0.316        Please use `context_schema` instead to specify the schema for run-scoped context.317 318 319    Example: Using entrypoint and tasks320        ```python321        import time322 323        from langgraph.func import entrypoint, task324        from langgraph.types import interrupt, Command325        from langgraph.checkpoint.memory import InMemorySaver326 327        @task328        def compose_essay(topic: str) -> str:329            time.sleep(1.0)  # Simulate slow operation330            return f"An essay about {topic}"331 332        @entrypoint(checkpointer=InMemorySaver())333        def review_workflow(topic: str) -> dict:334            \"\"\"Manages the workflow for generating and reviewing an essay.335 336            The workflow includes:337            1. Generating an essay about the given topic.338            2. Interrupting the workflow for human review of the generated essay.339 340            Upon resuming the workflow, compose_essay task will not be re-executed341            as its result is cached by the checkpointer.342 343            Args:344                topic: The subject of the essay.345 346            Returns:347                dict: A dictionary containing the generated essay and the human review.348            \"\"\"349            essay_future = compose_essay(topic)350            essay = essay_future.result()351            human_review = interrupt({352                \"question\": \"Please provide a review\",353                \"essay\": essay354            })355            return {356                \"essay\": essay,357                \"review\": human_review,358            }359 360        # Example configuration for the workflow361        config = {362            \"configurable\": {363                \"thread_id\": \"some_thread\"364            }365        }366 367        # Topic for the essay368        topic = \"cats\"369 370        # Stream the workflow to generate the essay and await human review371        for result in review_workflow.stream(topic, config):372            print(result)373 374        # Example human review provided after the interrupt375        human_review = \"This essay is great.\"376 377        # Resume the workflow with the provided human review378        for result in review_workflow.stream(Command(resume=human_review), config):379            print(result)380        ```381 382    Example: Accessing the previous return value383        When a checkpointer is enabled the function can access the previous return value384        of the previous invocation on the same thread id.385 386        ```python387        from typing import Optional388 389        from langgraph.checkpoint.memory import MemorySaver390 391        from langgraph.func import entrypoint392 393 394        @entrypoint(checkpointer=InMemorySaver())395        def my_workflow(input_data: str, previous: Optional[str] = None) -> str:396            return "world"397 398 399        config = {"configurable": {"thread_id": "some_thread"}}400        my_workflow.invoke("hello", config)401        ```402 403    Example: Using `entrypoint.final` to save a value404        The `entrypoint.final` object allows you to return a value while saving405        a different value to the checkpoint. This value will be accessible406        in the next invocation of the entrypoint via the `previous` parameter, as407        long as the same thread id is used.408 409        ```python410        from typing import Any411 412        from langgraph.checkpoint.memory import MemorySaver413 414        from langgraph.func import entrypoint415 416 417        @entrypoint(checkpointer=InMemorySaver())418        def my_workflow(419            number: int,420            *,421            previous: Any = None,422        ) -> entrypoint.final[int, int]:423            previous = previous or 0424            # This will return the previous value to the caller, saving425            # 2 * number to the checkpoint, which will be used in the next invocation426            # for the `previous` parameter.427            return entrypoint.final(value=previous, save=2 * number)428 429 430        config = {"configurable": {"thread_id": "some_thread"}}431 432        my_workflow.invoke(3, config)  # 0 (previous was None)433        my_workflow.invoke(1, config)  # 6 (previous was 3 * 2 from the previous invocation)434        ```435    """436 437    def __init__(438        self,439        checkpointer: BaseCheckpointSaver | None = None,440        store: BaseStore | None = None,441        cache: BaseCache | None = None,442        context_schema: type[ContextT] | None = None,443        cache_policy: CachePolicy | None = None,444        retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,445        timeout: float | timedelta | TimeoutPolicy | None = None,446        **kwargs: Unpack[DeprecatedKwargs],447    ) -> None:448        """Initialize the entrypoint decorator."""449        if (config_schema := kwargs.get("config_schema", MISSING)) is not MISSING:450            warnings.warn(451                "`config_schema` is deprecated and will be removed. Please use `context_schema` instead.",452                category=LangGraphDeprecatedSinceV10,453                stacklevel=2,454            )455            if context_schema is None:456                context_schema = cast(type[ContextT], config_schema)457 458        if (retry := kwargs.get("retry", MISSING)) is not MISSING:459            warnings.warn(460                "`retry` is deprecated and will be removed. Please use `retry_policy` instead.",461                category=LangGraphDeprecatedSinceV05,462                stacklevel=2,463            )464            if retry_policy is None:465                retry_policy = cast("RetryPolicy | Sequence[RetryPolicy]", retry)466 467        self.checkpointer = checkpointer468        self.store = store469        self.cache = cache470        self.cache_policy = cache_policy471        self.retry_policy = retry_policy472        self.timeout = coerce_timeout_policy(timeout)473        self.context_schema = context_schema474 475    @dataclass(**_DC_KWARGS)476    class final(Generic[R, S]):477        """A primitive that can be returned from an entrypoint.478 479        This primitive allows to save a value to the checkpointer distinct from the480        return value from the entrypoint.481 482        Example: Decoupling the return value and the save value483            ```python484            from langgraph.checkpoint.memory import InMemorySaver485            from langgraph.func import entrypoint486 487 488            @entrypoint(checkpointer=InMemorySaver())489            def my_workflow(490                number: int,491                *,492                previous: Any = None,493            ) -> entrypoint.final[int, int]:494                previous = previous or 0495                # This will return the previous value to the caller, saving496                # 2 * number to the checkpoint, which will be used in the next invocation497                # for the `previous` parameter.498                return entrypoint.final(value=previous, save=2 * number)499 500 501            config = {"configurable": {"thread_id": "1"}}502 503            my_workflow.invoke(3, config)  # 0 (previous was None)504            my_workflow.invoke(1, config)  # 6 (previous was 3 * 2 from the previous invocation)505            ```506        """507 508        value: R509        """Value to return. A value will always be returned even if it is `None`."""510        save: S511        """The value for the state for the next checkpoint.512 513        A value will always be saved even if it is `None`.514        """515 516    def __call__(self, func: Callable[..., Any]) -> Pregel:517        """Convert a function into a Pregel graph.518 519        Args:520            func: The function to convert. Support both sync and async functions.521 522        Returns:523            A Pregel graph.524        """525        # wrap generators in a function that writes to StreamWriter526        if inspect.isgeneratorfunction(func) or inspect.isasyncgenfunction(func):527            raise NotImplementedError(528                "Generators are not supported in the Functional API."529            )530 531        bound = get_runnable_for_entrypoint(func)532        stream_mode: StreamMode = "updates"533 534        # get input and output types535        sig = inspect.signature(func)536        first_parameter_name = next(iter(sig.parameters.keys()), None)537        if not first_parameter_name:538            raise ValueError("Entrypoint function must have at least one parameter")539        input_type = (540            sig.parameters[first_parameter_name].annotation541            if sig.parameters[first_parameter_name].annotation542            is not inspect.Signature.empty543            else Any544        )545 546        def _pluck_return_value(value: Any) -> Any:547            """Extract the return_ value the entrypoint.final object or passthrough."""548            return value.value if isinstance(value, entrypoint.final) else value549 550        def _pluck_save_value(value: Any) -> Any:551            """Get save value from the entrypoint.final object or passthrough."""552            return value.save if isinstance(value, entrypoint.final) else value553 554        output_type, save_type = Any, Any555        if sig.return_annotation is not inspect.Signature.empty:556            # User does not parameterize entrypoint.final properly557            if (558                sig.return_annotation is entrypoint.final559            ):  # Un-parameterized entrypoint.final560                output_type = save_type = Any561            else:562                origin = get_origin(sig.return_annotation)563                if origin is entrypoint.final:564                    type_annotations = get_args(sig.return_annotation)565                    if len(type_annotations) != 2:566                        raise TypeError(567                            "Please an annotation for both the return_ and "568                            "the save values."569                            "For example, `-> entrypoint.final[int, str]` would assign a "570                            "return_ a type of `int` and save the type `str`."571                        )572                    output_type, save_type = get_args(sig.return_annotation)573                else:574                    output_type = save_type = sig.return_annotation575 576        graph: Pregel[Any, ContextT, Any, Any] = Pregel(577            nodes={578                func.__name__: PregelNode(579                    bound=bound,580                    triggers=[START],581                    channels=START,582                    timeout=self.timeout,583                    writers=[584                        ChannelWrite(585                            [586                                ChannelWriteEntry(END, mapper=_pluck_return_value),587                                ChannelWriteEntry(PREVIOUS, mapper=_pluck_save_value),588                            ]589                        )590                    ],591                )592            },593            channels={594                START: EphemeralValue(input_type),595                END: LastValue(output_type, END),596                PREVIOUS: LastValue(save_type, PREVIOUS),597            },598            input_channels=START,599            output_channels=END,600            stream_channels=END,601            stream_mode=stream_mode,602            stream_eager=True,603            checkpointer=self.checkpointer,604            store=self.store,605            cache=self.cache,606            cache_policy=self.cache_policy,607            retry_policy=self.retry_policy or (),608            context_schema=self.context_schema,609        )610        if _serde.STRICT_MSGPACK_ENABLED:611            serde_allowlist = _serde.build_serde_allowlist(612                schemas=[input_type, output_type, save_type]613                + ([self.context_schema] if self.context_schema is not None else []),614                channels=graph.channels,615            )616            graph._serde_allowlist = serde_allowlist617            graph.checkpointer = _serde.apply_checkpointer_allowlist(618                graph.checkpointer, serde_allowlist619            )620        return graph621 
codekingpro/portable-devtools · Team Ai