Compare commits

...
Author SHA1 Message Date
William Fu-Hinthorn 396b3dc7f6 State modifier 2024-10-03 11:12:38 -07:00
vbarda 4fe00a138f lint 2024-10-03 12:17:40 -04:00
vbarda 1dcf33fc0e update 2024-10-03 12:16:46 -04:00
vbarda 97489f1386 langgraph: add support for passing store via state_modifier 2024-10-03 12:15:25 -04:00
3 changed files with 37 additions and 17 deletions
@@ -1,3 +1,4 @@
import inspect
from typing import Callable, Literal, Optional, Sequence, Type, TypeVar, Union, cast from typing import Callable, Literal, Optional, Sequence, Type, TypeVar, Union, cast
from langchain_core.language_models import BaseChatModel, LanguageModelLike from langchain_core.language_models import BaseChatModel, LanguageModelLike
@@ -20,6 +21,7 @@ from langgraph.prebuilt.tool_executor import ToolExecutor
from langgraph.prebuilt.tool_node import ToolNode from langgraph.prebuilt.tool_node import ToolNode
from langgraph.store.base import BaseStore from langgraph.store.base import BaseStore
from langgraph.types import Checkpointer from langgraph.types import Checkpointer
from langgraph.utils.runnable import RunnableCallable
# We create the AgentState that we will pass around # We create the AgentState that we will pass around
@@ -54,26 +56,29 @@ StateModifier = Union[
] ]
def _get_state_modifier_runnable(state_modifier: Optional[StateModifier]) -> Runnable: def _get_state_modifier_runnable(
state_modifier: Optional[StateModifier], store: Optional[BaseStore] = None
) -> Runnable:
state_modifier_runnable: Runnable state_modifier_runnable: Runnable
if state_modifier is None: if state_modifier is None:
state_modifier_runnable = RunnableLambda( state_modifier_runnable = RunnableLambda(
lambda state: state["messages"], name=STATE_MODIFIER_RUNNABLE_NAME lambda state, **kwargs: state["messages"], name=STATE_MODIFIER_RUNNABLE_NAME
) )
elif isinstance(state_modifier, str): elif isinstance(state_modifier, str):
_system_message: BaseMessage = SystemMessage(content=state_modifier) _system_message: BaseMessage = SystemMessage(content=state_modifier)
state_modifier_runnable = RunnableLambda( state_modifier_runnable = RunnableLambda(
lambda state: [_system_message] + state["messages"], lambda state, **kwargs: [_system_message] + state["messages"],
name=STATE_MODIFIER_RUNNABLE_NAME, name=STATE_MODIFIER_RUNNABLE_NAME,
) )
elif isinstance(state_modifier, SystemMessage): elif isinstance(state_modifier, SystemMessage):
state_modifier_runnable = RunnableLambda( state_modifier_runnable = RunnableLambda(
lambda state: [state_modifier] + state["messages"], lambda state, **kwargs: [state_modifier] + state["messages"],
name=STATE_MODIFIER_RUNNABLE_NAME, name=STATE_MODIFIER_RUNNABLE_NAME,
) )
elif callable(state_modifier): elif callable(state_modifier):
state_modifier_runnable = RunnableLambda( # Inspect the state_modifier signature
state_modifier, name=STATE_MODIFIER_RUNNABLE_NAME state_modifier_runnable = RunnableCallable(
state_modifier, name=STATE_MODIFIER_RUNNABLE_NAME, trace=True
) )
elif isinstance(state_modifier, Runnable): elif isinstance(state_modifier, Runnable):
state_modifier_runnable = state_modifier state_modifier_runnable = state_modifier
@@ -108,6 +113,7 @@ def _convert_messages_modifier_to_state_modifier(
def _get_model_preprocessing_runnable( def _get_model_preprocessing_runnable(
state_modifier: Optional[StateModifier], state_modifier: Optional[StateModifier],
messages_modifier: Optional[MessagesModifier], messages_modifier: Optional[MessagesModifier],
store: Optional[BaseStore],
) -> Runnable: ) -> Runnable:
# Add the state or message modifier, if exists # Add the state or message modifier, if exists
if state_modifier is not None and messages_modifier is not None: if state_modifier is not None and messages_modifier is not None:
@@ -118,7 +124,7 @@ def _get_model_preprocessing_runnable(
if state_modifier is None and messages_modifier is not None: if state_modifier is None and messages_modifier is not None:
state_modifier = _convert_messages_modifier_to_state_modifier(messages_modifier) state_modifier = _convert_messages_modifier_to_state_modifier(messages_modifier)
return _get_state_modifier_runnable(state_modifier) return _get_state_modifier_runnable(state_modifier, store)
def _should_bind_tools(model: LanguageModelLike, tools: Sequence[BaseTool]) -> bool: def _should_bind_tools(model: LanguageModelLike, tools: Sequence[BaseTool]) -> bool:
@@ -473,12 +479,20 @@ def create_react_agent(
else: else:
return "tools" return "tools"
preprocessor = _get_model_preprocessing_runnable(state_modifier, messages_modifier) # we're passing store here for validation
preprocessor = _get_model_preprocessing_runnable(
state_modifier, messages_modifier, store
)
model_runnable = preprocessor | model model_runnable = preprocessor | model
# Define the function that calls the model # Define the function that calls the model
def call_model(state: AgentState, config: RunnableConfig) -> AgentState: def call_model(
response = model_runnable.invoke(state, config) state: AgentState, config: RunnableConfig, *, store: BaseStore
) -> AgentState:
if store is not None:
response = model_runnable.invoke(state, config, store=store)
else:
response = model_runnable.invoke(state, config)
if ( if (
state["is_last_step"] state["is_last_step"]
and isinstance(response, AIMessage) and isinstance(response, AIMessage)
@@ -495,8 +509,13 @@ def create_react_agent(
# We return a list, because this will get added to the existing list # We return a list, because this will get added to the existing list
return {"messages": [response]} return {"messages": [response]}
async def acall_model(state: AgentState, config: RunnableConfig) -> AgentState: async def acall_model(
response = await model_runnable.ainvoke(state, config) state: AgentState, config: RunnableConfig, *, store: BaseStore
) -> AgentState:
if store is not None:
response = await model_runnable.ainvoke(state, config, store=store)
else:
response = await model_runnable.ainvoke(state, config)
if ( if (
state["is_last_step"] state["is_last_step"]
and isinstance(response, AIMessage) and isinstance(response, AIMessage)
@@ -517,7 +536,7 @@ def create_react_agent(
workflow = StateGraph(state_schema or AgentState) workflow = StateGraph(state_schema or AgentState)
# Define the two nodes we will cycle between # Define the two nodes we will cycle between
workflow.add_node("agent", RunnableLambda(call_model, acall_model)) workflow.add_node("agent", RunnableCallable(call_model, acall_model))
workflow.add_node("tools", tool_node) workflow.add_node("tools", tool_node)
# Set the entrypoint as `agent` # Set the entrypoint as `agent`
@@ -159,6 +159,7 @@ class RunnableCallable(Runnable):
) )
elif kwargs.get(kw) is None: elif kwargs.get(kw) is None:
kwargs[kw] = _conf.get(ck, defv) kwargs[kw] = _conf.get(ck, defv)
context = copy_context() context = copy_context()
if self.trace: if self.trace:
callback_manager = get_callback_manager_for_config(config, self.tags) callback_manager = get_callback_manager_for_config(config, self.tags)
@@ -4841,10 +4841,10 @@
"type": "runnable", "type": "runnable",
"data": { "data": {
"id": [ "id": [
"langchain_core", "langgraph",
"runnables", "utils",
"base", "runnable",
"RunnableLambda" "RunnableCallable"
], ],
"name": "agent" "name": "agent"
} }