Merge pull request #2679 from langchain-ai/nc/9dec/imperative-generator

lib: imperative api: Generators use yield to publish stream_mode=custom events
This commit is contained in:
Nuno Campos
2024-12-09 08:22:31 -08:00
committed by GitHub
+26 -3
View File
@@ -1,6 +1,7 @@
import asyncio import asyncio
import concurrent import concurrent
import concurrent.futures import concurrent.futures
import inspect
import types import types
from functools import partial, update_wrapper from functools import partial, update_wrapper
from typing import ( from typing import (
@@ -24,7 +25,7 @@ from langgraph.pregel.call import get_runnable_for_func
from langgraph.pregel.read import PregelNode from langgraph.pregel.read import PregelNode
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
from langgraph.store.base import BaseStore from langgraph.store.base import BaseStore
from langgraph.types import RetryPolicy from langgraph.types import RetryPolicy, StreamMode, StreamWriter
P = ParamSpec("P") P = ParamSpec("P")
P1 = TypeVar("P1") P1 = TypeVar("P1")
@@ -76,10 +77,32 @@ def entrypoint(
store: Optional[BaseStore] = None, store: Optional[BaseStore] = None,
) -> Callable[[types.FunctionType], Pregel]: ) -> Callable[[types.FunctionType], Pregel]:
def _imp(func: types.FunctionType) -> Pregel: def _imp(func: types.FunctionType) -> Pregel:
if inspect.isgeneratorfunction(func):
def gen_wrapper(*args: Any, writer: StreamWriter, **kwargs: Any) -> Any:
for chunk in func(*args, **kwargs):
writer(chunk)
bound = get_runnable_for_func(gen_wrapper)
stream_mode: StreamMode = "custom"
elif inspect.isasyncgenfunction(func):
async def agen_wrapper(
*args: Any, writer: StreamWriter, **kwargs: Any
) -> Any:
async for chunk in func(*args, **kwargs):
writer(chunk)
bound = get_runnable_for_func(agen_wrapper)
stream_mode = "custom"
else:
bound = get_runnable_for_func(func)
stream_mode = "updates"
return Pregel( return Pregel(
nodes={ nodes={
func.__name__: PregelNode( func.__name__: PregelNode(
bound=get_runnable_for_func(func), bound=bound,
triggers=[START], triggers=[START],
channels=[START], channels=[START],
writers=[ChannelWrite([ChannelWriteEntry(END)], tags=[TAG_HIDDEN])], writers=[ChannelWrite([ChannelWriteEntry(END)], tags=[TAG_HIDDEN])],
@@ -89,7 +112,7 @@ def entrypoint(
input_channels=START, input_channels=START,
output_channels=END, output_channels=END,
stream_channels=END, stream_channels=END,
stream_mode="updates", stream_mode=stream_mode,
checkpointer=checkpointer, checkpointer=checkpointer,
store=store, store=store,
) )