For stream_mode=messages skip any nodes/llms with tag nostream

This commit is contained in:
Nuno Campos
2024-10-17 15:28:14 -07:00
parent 4dfdb9a83e
commit e2a3698250
2 changed files with 9 additions and 3 deletions
+2
View File
@@ -12,6 +12,8 @@ EMPTY_MAP: Mapping[str, Any] = MappingProxyType({})
EMPTY_SEQ: tuple[str, ...] = tuple()
# --- Public constants ---
TAG_NOSTREAM = sys.intern("langsmith:nostream")
"""Tag to disable streaming for a chat model."""
TAG_HIDDEN = sys.intern("langsmith:hidden")
"""Tag to hide a node/edge from certain tracing/streaming environments."""
START = sys.intern("__start__")
+7 -3
View File
@@ -17,7 +17,7 @@ from langchain_core.messages import BaseMessage
from langchain_core.outputs import ChatGenerationChunk, LLMResult
from langchain_core.tracers._streaming import T, _StreamingCallbackHandler
from langgraph.constants import NS_SEP
from langgraph.constants import NS_SEP, TAG_HIDDEN, TAG_NOSTREAM
from langgraph.pregel.loop import StreamChunk
Meta = tuple[tuple[str, ...], dict[str, Any]]
@@ -63,7 +63,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
metadata: Optional[dict[str, Any]] = None,
**kwargs: Any,
) -> Any:
if metadata:
if metadata and (not tags or TAG_NOSTREAM not in tags):
self.metadata[run_id] = (
tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP)),
metadata,
@@ -114,7 +114,11 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
metadata: Optional[Dict[str, Any]] = None,
**kwargs: Any,
) -> Any:
if metadata and kwargs.get("name") == metadata.get("langgraph_node"):
if (
metadata
and kwargs.get("name") == metadata.get("langgraph_node")
and (not tags or TAG_HIDDEN not in tags)
):
self.metadata[run_id] = (
tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP)),
metadata,