Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_runnable.py943 linesDownload Raw Back to _internal
1from __future__ import annotations2 3import asyncio4import enum5import inspect6import sys7import warnings8from collections.abc import (9    AsyncIterator,10    Awaitable,11    Callable,12    Coroutine,13    Generator,14    Iterator,15    Sequence,16)17from contextlib import AsyncExitStack, contextmanager18from contextvars import Context, Token, copy_context19from functools import partial, wraps20from typing import (21    Any,22    Optional,23    Protocol,24    TypeGuard,25    cast,26)27 28from langchain_core.runnables.base import (29    Runnable,30    RunnableConfig,31    RunnableLambda,32    RunnableParallel,33    RunnableSequence,34)35from langchain_core.runnables.base import (36    RunnableLike as LCRunnableLike,37)38from langchain_core.runnables.config import (39    run_in_executor,40    var_child_runnable_config,41)42from langchain_core.runnables.utils import Input, Output43from langchain_core.tracers.langchain import LangChainTracer44from langgraph.store.base import BaseStore45 46from langgraph._internal._config import (47    ensure_config,48    get_async_callback_manager_for_config,49    get_callback_manager_for_config,50    patch_config,51)52from langgraph._internal._constants import (53    CONF,54    CONFIG_KEY_NODE_ERROR,55    CONFIG_KEY_RUNTIME,56)57from langgraph._internal._typing import MISSING58from langgraph.errors import NodeError59from langgraph.types import StreamWriter60 61try:62    from langchain_core.tracers._streaming import _StreamingCallbackHandler63except ImportError:64    _StreamingCallbackHandler = None  # type: ignore65 66 67def _set_config_context(68    config: RunnableConfig, run: Any = None69) -> Token[RunnableConfig | None]:70    """Set the child Runnable config + tracing context.71 72    Args:73        config: The config to set.74    """75    config_token = var_child_runnable_config.set(config)76    if run is not None:77        from langsmith.run_helpers import _set_tracing_context78 79        _set_tracing_context({"parent": run})80    return config_token81 82 83def _unset_config_context(token: Token[RunnableConfig | None], run: Any = None) -> None:84    """Set the child Runnable config + tracing context.85 86    Args:87        token: The config token to reset.88    """89    var_child_runnable_config.reset(token)90    if run is not None:91        from langsmith.run_helpers import _set_tracing_context92 93        _set_tracing_context(94            {95                "parent": None,96                "project_name": None,97                "tags": None,98                "metadata": None,99                "enabled": None,100                "client": None,101            }102        )103 104 105@contextmanager106def set_config_context(107    config: RunnableConfig, run: Any = None108) -> Generator[Context, None, None]:109    """Set the child Runnable config + tracing context.110 111    Args:112        config: The config to set.113    """114    ctx = copy_context()115    config_token = ctx.run(_set_config_context, config, run)116    try:117        yield ctx118    finally:119        ctx.run(_unset_config_context, config_token, run)120 121 122def create_task_in_config_context(123    coro_factory: Callable[[], Coroutine[Any, Any, Any]], config: RunnableConfig124) -> asyncio.Task[Any]:125    """Create an asyncio.Task that inherits `config` as the child runnable context.126 127    `asyncio.create_task` snapshots the current contextvars onto the new task,128    so calling `create_task` while the config context is set ensures the task129    sees `config` via `var_child_runnable_config` and any tracing parent.130    """131    with set_config_context(config) as context:132        return context.run(lambda: asyncio.create_task(coro_factory()))133 134 135# Before Python 3.11 native StrEnum is not available136class StrEnum(str, enum.Enum):137    """A string enum."""138 139 140# Special type to denote any type is accepted141ANY_TYPE = object()142 143ASYNCIO_ACCEPTS_CONTEXT = sys.version_info >= (3, 11)144 145# List of keyword arguments that can be injected into nodes / tasks / tools at runtime.146# A named argument may appear multiple times if it appears with distinct types.147KWARGS_CONFIG_KEYS: tuple[tuple[str, tuple[Any, ...], str, Any], ...] = (148    (149        "config",150        (151            RunnableConfig,152            "RunnableConfig",153            Optional[RunnableConfig],  # noqa: UP045154            "Optional[RunnableConfig]",155            inspect.Parameter.empty,156        ),157        # for now, use config directly, eventually, will pop off of Runtime158        "N/A",159        inspect.Parameter.empty,160    ),161    (162        "writer",163        (StreamWriter, "StreamWriter", inspect.Parameter.empty),164        "stream_writer",165        lambda _: None,166    ),167    (168        "store",169        (170            BaseStore,171            "BaseStore",172            inspect.Parameter.empty,173        ),174        "store",175        inspect.Parameter.empty,176    ),177    (178        "store",179        (180            Optional[BaseStore],  # noqa: UP045181            "Optional[BaseStore]",182        ),183        "store",184        None,185    ),186    (187        "previous",188        (ANY_TYPE,),189        "previous",190        inspect.Parameter.empty,191    ),192    (193        "runtime",194        (ANY_TYPE,),195        # we never hit this block, we just inject runtime directly196        "N/A",197        inspect.Parameter.empty,198    ),199    (200        "error",201        (NodeError, "NodeError"),202        # we never hit this block, we read directly from configurable203        "N/A",204        # default to None so non-handler nodes that happen to type a parameter205        # `error: NodeError` don't blow up; handlers always receive a NodeError.206        None,207    ),208)209"""List of kwargs that can be passed to functions, and their corresponding210config keys, default values and type annotations.211 212Used to configure keyword arguments that can be injected at runtime213from the `Runtime` object as kwargs to `invoke`, `ainvoke`, `stream` and `astream`.214 215For a keyword to be injected from the config object, the function signature216must contain a kwarg with the same name and a matching type annotation.217 218Each tuple contains:219- the name of the kwarg in the function signature220- the type annotation(s) for the kwarg221- the `Runtime` attribute for fetching the value (N/A if not applicable)222 223This is fully internal and should be further refactored to use `get_type_hints`224to resolve forward references and optional types formatted like BaseStore | None.225"""226 227VALID_KINDS = (inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.KEYWORD_ONLY)228 229 230class _RunnableWithWriter(Protocol[Input, Output]):231    def __call__(self, state: Input, *, writer: StreamWriter) -> Output: ...232 233 234class _RunnableWithStore(Protocol[Input, Output]):235    def __call__(self, state: Input, *, store: BaseStore) -> Output: ...236 237 238class _RunnableWithWriterStore(Protocol[Input, Output]):239    def __call__(240        self, state: Input, *, writer: StreamWriter, store: BaseStore241    ) -> Output: ...242 243 244class _RunnableWithConfigWriter(Protocol[Input, Output]):245    def __call__(246        self, state: Input, *, config: RunnableConfig, writer: StreamWriter247    ) -> Output: ...248 249 250class _RunnableWithConfigStore(Protocol[Input, Output]):251    def __call__(252        self, state: Input, *, config: RunnableConfig, store: BaseStore253    ) -> Output: ...254 255 256class _RunnableWithConfigWriterStore(Protocol[Input, Output]):257    def __call__(258        self,259        state: Input,260        *,261        config: RunnableConfig,262        writer: StreamWriter,263        store: BaseStore,264    ) -> Output: ...265 266 267RunnableLike = (268    LCRunnableLike269    | _RunnableWithWriter[Input, Output]270    | _RunnableWithStore[Input, Output]271    | _RunnableWithWriterStore[Input, Output]272    | _RunnableWithConfigWriter[Input, Output]273    | _RunnableWithConfigStore[Input, Output]274    | _RunnableWithConfigWriterStore[Input, Output]275)276 277 278class RunnableCallable(Runnable):279    """A much simpler version of RunnableLambda that requires sync and async functions."""280 281    def __init__(282        self,283        func: Callable[..., Any | Runnable] | None,284        afunc: Callable[..., Awaitable[Any | Runnable]] | None = None,285        *,286        name: str | None = None,287        tags: Sequence[str] | None = None,288        trace: bool = True,289        recurse: bool = True,290        explode_args: bool = False,291        **kwargs: Any,292    ) -> None:293        self.name = name294        if self.name is None:295            if func:296                try:297                    if func.__name__ != "<lambda>":298                        self.name = func.__name__299                except AttributeError:300                    pass301            elif afunc:302                try:303                    self.name = afunc.__name__304                except AttributeError:305                    pass306        self.func = func307        self.afunc = afunc308        self.tags = tags309        self.kwargs = kwargs310        self.trace = trace311        self.recurse = recurse312        self.explode_args = explode_args313        # check signature314        if func is None and afunc is None:315            raise ValueError("At least one of func or afunc must be provided.")316 317        self.func_accepts: dict[str, tuple[str, Any]] = {}318        params = inspect.signature(cast(Callable, func or afunc)).parameters319 320        for kw, typ, runtime_key, default in KWARGS_CONFIG_KEYS:321            p = params.get(kw)322 323            if p is None or p.kind not in VALID_KINDS:324                # If parameter is not found or is not a valid kind, skip325                continue326 327            if typ != (ANY_TYPE,) and p.annotation not in typ:328                # A specific type is required, but the function annotation does329                # not match the expected type.330 331                # If this is a config parameter with incorrect typing, emit a warning332                # because we used to support any type but are moving towards more correct typing333                if kw == "config" and p.annotation != inspect.Parameter.empty:334                    warnings.warn(335                        f"The 'config' parameter should be typed as 'RunnableConfig' or "336                        f"'RunnableConfig | None', not '{p.annotation}'. ",337                        UserWarning,338                        stacklevel=4,339                    )340                continue341 342            # If the kwarg is accepted by the function, store the key / runtime attribute to inject343            self.func_accepts[kw] = (runtime_key, default)344 345    def __repr__(self) -> str:346        repr_args = {347            k: v348            for k, v in self.__dict__.items()349            if k not in {"name", "func", "afunc", "config", "kwargs", "trace"}350        }351        return f"{self.get_name()}({', '.join(f'{k}={v!r}' for k, v in repr_args.items())})"352 353    def invoke(354        self, input: Any, config: RunnableConfig | None = None, **kwargs: Any355    ) -> Any:356        if self.func is None:357            raise TypeError(358                f'No synchronous function provided to "{self.name}".'359                "\nEither initialize with a synchronous function or invoke"360                " via the async API (ainvoke, astream, etc.)"361            )362        if config is None:363            config = ensure_config()364        if self.explode_args:365            args, _kwargs = input366            kwargs = {**self.kwargs, **_kwargs, **kwargs}367        else:368            args = (input,)369            kwargs = {**self.kwargs, **kwargs}370 371        runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME)372 373        for kw, (runtime_key, default) in self.func_accepts.items():374            # If the kwarg is already set, use the set value375            if kw in kwargs:376                continue377 378            kw_value: Any = MISSING379            if kw == "config":380                kw_value = config381            elif kw == "error":382                kw_value = config.get(CONF, {}).get(CONFIG_KEY_NODE_ERROR, MISSING)383            elif runtime:384                if kw == "runtime":385                    kw_value = runtime386                else:387                    try:388                        kw_value = getattr(runtime, runtime_key)389                    except AttributeError:390                        pass391 392            if kw_value is MISSING:393                if default is inspect.Parameter.empty:394                    raise ValueError(395                        f"Missing required config key '{runtime_key}' for '{self.name}'."396                    )397                kw_value = default398            kwargs[kw] = kw_value399 400        if self.trace:401            callback_manager = get_callback_manager_for_config(config, self.tags)402            run_manager = callback_manager.on_chain_start(403                None,404                input,405                name=config.get("run_name") or self.get_name(),406                run_id=config.pop("run_id", None),407            )408            try:409                child_config = patch_config(config, callbacks=run_manager.get_child())410                # get the run411                for h in run_manager.handlers:412                    if isinstance(h, LangChainTracer):413                        run = h.run_map.get(str(run_manager.run_id))414                        break415                else:416                    run = None417                # run in context418                with set_config_context(child_config, run) as context:419                    ret = context.run(self.func, *args, **kwargs)420            except BaseException as e:421                run_manager.on_chain_error(e)422                raise423            else:424                run_manager.on_chain_end(ret)425        else:426            ret = self.func(*args, **kwargs)427        if self.recurse and isinstance(ret, Runnable):428            return ret.invoke(input, config)429        return ret430 431    async def ainvoke(432        self, input: Any, config: RunnableConfig | None = None, **kwargs: Any433    ) -> Any:434        if not self.afunc:435            return self.invoke(input, config)436        if config is None:437            config = ensure_config()438        if self.explode_args:439            args, _kwargs = input440            kwargs = {**self.kwargs, **_kwargs, **kwargs}441        else:442            args = (input,)443            kwargs = {**self.kwargs, **kwargs}444 445        runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME)446 447        for kw, (runtime_key, default) in self.func_accepts.items():448            # If the kwarg has already been set, use the set value449            if kw in kwargs:450                continue451 452            kw_value: Any = MISSING453            if kw == "config":454                kw_value = config455            elif kw == "error":456                kw_value = config.get(CONF, {}).get(CONFIG_KEY_NODE_ERROR, MISSING)457            elif runtime:458                if kw == "runtime":459                    kw_value = runtime460                else:461                    try:462                        kw_value = getattr(runtime, runtime_key)463                    except AttributeError:464                        pass465            if kw_value is MISSING:466                if default is inspect.Parameter.empty:467                    raise ValueError(468                        f"Missing required config key '{runtime_key}' for '{self.name}'."469                    )470                kw_value = default471            kwargs[kw] = kw_value472 473        if self.trace:474            callback_manager = get_async_callback_manager_for_config(config, self.tags)475            run_manager = await callback_manager.on_chain_start(476                None,477                input,478                name=config.get("run_name") or self.name,479                run_id=config.pop("run_id", None),480            )481            try:482                child_config = patch_config(config, callbacks=run_manager.get_child())483                coro = cast(Coroutine[None, None, Any], self.afunc(*args, **kwargs))484                if ASYNCIO_ACCEPTS_CONTEXT:485                    for h in run_manager.handlers:486                        if isinstance(h, LangChainTracer):487                            run = h.run_map.get(str(run_manager.run_id))488                            break489                    else:490                        run = None491                    with set_config_context(child_config, run) as context:492                        ret = await asyncio.create_task(coro, context=context)493                else:494                    ret = await coro495            except BaseException as e:496                await run_manager.on_chain_error(e)497                raise498            else:499                await run_manager.on_chain_end(ret)500        else:501            ret = await self.afunc(*args, **kwargs)502        if self.recurse and isinstance(ret, Runnable):503            return await ret.ainvoke(input, config)504        return ret505 506 507def is_async_callable(508    func: Any,509) -> TypeGuard[Callable[..., Awaitable]]:510    """Check if a function is async."""511    return (512        inspect.iscoroutinefunction(func)513        or hasattr(func, "__call__")514        and inspect.iscoroutinefunction(func.__call__)515    )516 517 518def is_async_generator(519    func: Any,520) -> TypeGuard[Callable[..., AsyncIterator]]:521    """Check if a function is an async generator."""522    return (523        inspect.isasyncgenfunction(func)524        or hasattr(func, "__call__")525        and inspect.isasyncgenfunction(func.__call__)526    )527 528 529def coerce_to_runnable(530    thing: RunnableLike, *, name: str | None, trace: bool531) -> Runnable:532    """Coerce a runnable-like object into a Runnable.533 534    Args:535        thing: A runnable-like object.536 537    Returns:538        A Runnable.539    """540    if isinstance(thing, Runnable):541        return thing542    elif is_async_generator(thing) or inspect.isgeneratorfunction(thing):543        return RunnableLambda(thing, name=name)544    elif callable(thing):545        if is_async_callable(thing):546            return RunnableCallable(None, thing, name=name, trace=trace)547        else:548            return RunnableCallable(549                thing,550                wraps(thing)(partial(run_in_executor, None, thing)),  # type: ignore[arg-type]551                name=name,552                trace=trace,553            )554    elif isinstance(thing, dict):555        return RunnableParallel(thing)556    else:557        raise TypeError(558            f"Expected a Runnable, callable or dict."559            f"Instead got an unsupported type: {type(thing)}"560        )561 562 563class RunnableSeq(Runnable):564    """Sequence of `Runnable`, where the output of each is the input of the next.565 566    `RunnableSeq` is a simpler version of `RunnableSequence` that is internal to567    LangGraph.568    """569 570    def __init__(571        self,572        *steps: RunnableLike,573        name: str | None = None,574        trace_inputs: Callable[[Any], Any] | None = None,575    ) -> None:576        """Create a new RunnableSeq.577 578        Args:579            steps: The steps to include in the sequence.580            name: The name of the `Runnable`.581 582        Raises:583            ValueError: If the sequence has less than 2 steps.584        """585        steps_flat: list[Runnable] = []586        for step in steps:587            if isinstance(step, RunnableSequence):588                steps_flat.extend(step.steps)589            elif isinstance(step, RunnableSeq):590                steps_flat.extend(step.steps)591            else:592                steps_flat.append(coerce_to_runnable(step, name=None, trace=True))593        if len(steps_flat) < 2:594            raise ValueError(595                f"RunnableSeq must have at least 2 steps, got {len(steps_flat)}"596            )597        self.steps = steps_flat598        self.name = name599        self.trace_inputs = trace_inputs600 601    def __or__(602        self,603        other: Any,604    ) -> Runnable:605        if isinstance(other, RunnableSequence):606            return RunnableSeq(607                *self.steps,608                other.first,609                *other.middle,610                other.last,611                name=self.name or other.name,612            )613        elif isinstance(other, RunnableSeq):614            return RunnableSeq(615                *self.steps,616                *other.steps,617                name=self.name or other.name,618            )619        else:620            return RunnableSeq(621                *self.steps,622                coerce_to_runnable(other, name=None, trace=True),623                name=self.name,624            )625 626    def __ror__(627        self,628        other: Any,629    ) -> Runnable:630        if isinstance(other, RunnableSequence):631            return RunnableSequence(632                other.first,633                *other.middle,634                other.last,635                *self.steps,636                name=other.name or self.name,637            )638        elif isinstance(other, RunnableSeq):639            return RunnableSeq(640                *other.steps,641                *self.steps,642                name=other.name or self.name,643            )644        else:645            return RunnableSequence(646                coerce_to_runnable(other, name=None, trace=True),647                *self.steps,648                name=self.name,649            )650 651    def invoke(652        self, input: Input, config: RunnableConfig | None = None, **kwargs: Any653    ) -> Any:654        if config is None:655            config = ensure_config()656        # setup callbacks and context657        callback_manager = get_callback_manager_for_config(config)658        # start the root run659        run_manager = callback_manager.on_chain_start(660            None,661            self.trace_inputs(input) if self.trace_inputs is not None else input,662            name=config.get("run_name") or self.get_name(),663            run_id=config.pop("run_id", None),664        )665        # invoke all steps in sequence666        try:667            for i, step in enumerate(self.steps):668                # mark each step as a child run669                config = patch_config(670                    config, callbacks=run_manager.get_child(f"seq:step:{i + 1}")671                )672                # 1st step is the actual node,673                # others are writers which don't need to be run in context674                if i == 0:675                    # get the run object676                    for h in run_manager.handlers:677                        if isinstance(h, LangChainTracer):678                            run = h.run_map.get(str(run_manager.run_id))679                            break680                    else:681                        run = None682                    # run in context683                    with set_config_context(config, run) as context:684                        input = context.run(step.invoke, input, config, **kwargs)685                else:686                    input = step.invoke(input, config)687        # finish the root run688        except BaseException as e:689            run_manager.on_chain_error(e)690            raise691        else:692            run_manager.on_chain_end(input)693            return input694 695    async def ainvoke(696        self,697        input: Input,698        config: RunnableConfig | None = None,699        **kwargs: Any | None,700    ) -> Any:701        if config is None:702            config = ensure_config()703        # setup callbacks704        callback_manager = get_async_callback_manager_for_config(config)705        # start the root run706        run_manager = await callback_manager.on_chain_start(707            None,708            self.trace_inputs(input) if self.trace_inputs is not None else input,709            name=config.get("run_name") or self.get_name(),710            run_id=config.pop("run_id", None),711        )712 713        # invoke all steps in sequence714        try:715            for i, step in enumerate(self.steps):716                # mark each step as a child run717                config = patch_config(718                    config, callbacks=run_manager.get_child(f"seq:step:{i + 1}")719                )720                # 1st step is the actual node,721                # others are writers which don't need to be run in context722                if i == 0:723                    if ASYNCIO_ACCEPTS_CONTEXT:724                        # get the run object725                        for h in run_manager.handlers:726                            if isinstance(h, LangChainTracer):727                                run = h.run_map.get(str(run_manager.run_id))728                                break729                        else:730                            run = None731                        # run in context732                        with set_config_context(config, run) as context:733                            input = await asyncio.create_task(734                                step.ainvoke(input, config, **kwargs), context=context735                            )736                    else:737                        input = await step.ainvoke(input, config, **kwargs)738                else:739                    input = await step.ainvoke(input, config)740        # finish the root run741        except BaseException as e:742            await run_manager.on_chain_error(e)743            raise744        else:745            await run_manager.on_chain_end(input)746            return input747 748    def stream(749        self,750        input: Input,751        config: RunnableConfig | None = None,752        **kwargs: Any | None,753    ) -> Iterator[Any]:754        if config is None:755            config = ensure_config()756        # setup callbacks757        callback_manager = get_callback_manager_for_config(config)758        # start the root run759        run_manager = callback_manager.on_chain_start(760            None,761            self.trace_inputs(input) if self.trace_inputs is not None else input,762            name=config.get("run_name") or self.get_name(),763            run_id=config.pop("run_id", None),764        )765        # get the run object766        for h in run_manager.handlers:767            if isinstance(h, LangChainTracer):768                run = h.run_map.get(str(run_manager.run_id))769                break770        else:771            run = None772        # create first step config773        config = patch_config(774            config,775            callbacks=run_manager.get_child(f"seq:step:{1}"),776        )777        # run all in context778        with set_config_context(config, run) as context:779            try:780                # stream the last steps781                # transform the input stream of each step with the next782                # steps that don't natively support transforming an input stream will783                # buffer input in memory until all available, and then start emitting output784                for idx, step in enumerate(self.steps):785                    if idx == 0:786                        iterator = step.stream(input, config, **kwargs)787                    else:788                        config = patch_config(789                            config,790                            callbacks=run_manager.get_child(f"seq:step:{idx + 1}"),791                        )792                        iterator = step.transform(iterator, config)793                # populates streamed_output in astream_log() output if needed794                if _StreamingCallbackHandler is not None:795                    for h in run_manager.handlers:796                        if isinstance(h, _StreamingCallbackHandler):797                            iterator = h.tap_output_iter(run_manager.run_id, iterator)798                # consume into final output799                output = context.run(_consume_iter, iterator)800                # sequence doesn't emit output, yield to mark as generator801                yield802            except BaseException as e:803                run_manager.on_chain_error(e)804                raise805            else:806                run_manager.on_chain_end(output)807 808    async def astream(809        self,810        input: Input,811        config: RunnableConfig | None = None,812        **kwargs: Any | None,813    ) -> AsyncIterator[Any]:814        if config is None:815            config = ensure_config()816        # setup callbacks817        callback_manager = get_async_callback_manager_for_config(config)818        # start the root run819        run_manager = await callback_manager.on_chain_start(820            None,821            self.trace_inputs(input) if self.trace_inputs is not None else input,822            name=config.get("run_name") or self.get_name(),823            run_id=config.pop("run_id", None),824        )825        # stream the last steps826        # transform the input stream of each step with the next827        # steps that don't natively support transforming an input stream will828        # buffer input in memory until all available, and then start emitting output829        if ASYNCIO_ACCEPTS_CONTEXT:830            # get the run object831            for h in run_manager.handlers:832                if isinstance(h, LangChainTracer):833                    run = h.run_map.get(str(run_manager.run_id))834                    break835            else:836                run = None837            # create first step config838            config = patch_config(839                config,840                callbacks=run_manager.get_child(f"seq:step:{1}"),841            )842            # run all in context843            with set_config_context(config, run) as context:844                try:845                    async with AsyncExitStack() as stack:846                        for idx, step in enumerate(self.steps):847                            if idx == 0:848                                aiterator = step.astream(input, config, **kwargs)849                            else:850                                config = patch_config(851                                    config,852                                    callbacks=run_manager.get_child(853                                        f"seq:step:{idx + 1}"854                                    ),855                                )856                                aiterator = step.atransform(aiterator, config)857                            if hasattr(aiterator, "aclose"):858                                stack.push_async_callback(aiterator.aclose)859                        # populates streamed_output in astream_log() output if needed860                        if _StreamingCallbackHandler is not None:861                            for h in run_manager.handlers:862                                if isinstance(h, _StreamingCallbackHandler):863                                    aiterator = h.tap_output_aiter(864                                        run_manager.run_id, aiterator865                                    )866                        # consume into final output867                        output = await asyncio.create_task(868                            _consume_aiter(aiterator), context=context869                        )870                        # sequence doesn't emit output, yield to mark as generator871                        yield872                except BaseException as e:873                    await run_manager.on_chain_error(e)874                    raise875                else:876                    await run_manager.on_chain_end(output)877        else:878            try:879                async with AsyncExitStack() as stack:880                    for idx, step in enumerate(self.steps):881                        config = patch_config(882                            config,883                            callbacks=run_manager.get_child(f"seq:step:{idx + 1}"),884                        )885                        if idx == 0:886                            aiterator = step.astream(input, config, **kwargs)887                        else:888                            aiterator = step.atransform(aiterator, config)889                        if hasattr(aiterator, "aclose"):890                            stack.push_async_callback(aiterator.aclose)891                    # populates streamed_output in astream_log() output if needed892                    if _StreamingCallbackHandler is not None:893                        for h in run_manager.handlers:894                            if isinstance(h, _StreamingCallbackHandler):895                                aiterator = h.tap_output_aiter(896                                    run_manager.run_id, aiterator897                                )898                    # consume into final output899                    output = await _consume_aiter(aiterator)900                    # sequence doesn't emit output, yield to mark as generator901                    yield902            except BaseException as e:903                await run_manager.on_chain_error(e)904                raise905            else:906                await run_manager.on_chain_end(output)907 908 909def _consume_iter(it: Iterator[Any]) -> Any:910    """Consume an iterator."""911    output: Any = None912    add_supported = False913    for chunk in it:914        # collect final output915        if output is None:916            output = chunk917        elif add_supported:918            try:919                output = output + chunk920            except TypeError:921                output = chunk922                add_supported = False923        else:924            output = chunk925    return output926 927 928async def _consume_aiter(it: AsyncIterator[Any]) -> Any:929    """Consume an async iterator."""930    output: Any = None931    add_supported = False932    async for chunk in it:933        # collect final output934        if add_supported:935            try:936                output = output + chunk937            except TypeError:938                output = chunk939                add_supported = False940        else:941            output = chunk942    return output943 
codekingpro/portable-devtools · Team Ai