mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-27 12:04:58 +02:00
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:
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user