from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
from langchain_core.runnables import Runnable
import sys
from langchain.agents import create_agent
from langchain.tools import ToolRuntime, BaseTool
from langgraph.graph import StateGraph, START, END
from langchain_core.messages import HumanMessage, AIMessage, ToolCall
from langchain_core.runnables import Runnable
from langchain.agents.middleware import wrap_model_call, ModelRequest, ModelResponse
from langchain_core.tools import tool
from langgraph.graph import add_messages
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.types import interrupt, Command
from typing import Annotated, Any, Awaitable, Callable, Literal, Optional
from typing_extensions import TypedDict
import structlog
from pydantic import BaseModel
from enum import Enum
import uuid
class FakeChatModel(FakeMessagesListChatModel):
def bind_tools(self, tools, *, tool_choice = None, **kwargs) -> Runnable:
return super().bind(**kwargs)
checkpoint = InMemorySaver()
class AgentState(TypedDict):
messages: Annotated[list, add_messages]
class InterruptType(str, Enum):
AGUI_TOOL_CALL= "agui-tool-call"
class BaseInterrupt(BaseModel):
type: InterruptType
class ToolCallInterrupt(BaseInterrupt):
type: Literal[InterruptType.AGUI_TOOL_CALL]
tool_call_id: str
tool_call_name: str
tool_call_args: dict
class BaseResume(BaseModel):
type: InterruptType
class ToolCallResume(BaseResume):
type: Literal[InterruptType.AGUI_TOOL_CALL]
tool_call_id: str
content: str
structlog.configure(logger_factory=structlog.PrintLoggerFactory(sys.stderr))
logger = structlog.get_logger()
weather_schema = {
"type": "object",
"properties": {
"location": {"type": "string"},
"units": {"type": "string"},
"include_forecast": {"type": "boolean"},
},
"required": ["location", "units", "include_forecast"]
}
@tool(args_schema=weather_schema)
async def get_weather(runtime: ToolRuntime | None = None, **kwargs) -> str:
"""Get current weather and optional forecast."""
await logger.ainfo("get_weather")
frontend_response=interrupt(
ToolCallInterrupt(
type=InterruptType.AGUI_TOOL_CALL,
tool_call_id=runtime.tool_call_id,
tool_call_name="get_weather",
tool_call_args=kwargs,
))
await logger.ainfo("get_weather", resp=frontend_response)
return frontend_response.content
@tool
async def say_hello(name: str, runtime: ToolRuntime | None = None, **kwargs) -> str:
"""Say hello."""
resp = interrupt(
ToolCallInterrupt(
type=InterruptType.AGUI_TOOL_CALL,
tool_call_id=runtime.tool_call_id,
tool_call_name="say_hello",
tool_call_args=kwargs,
)
)
return resp.content
def create_weather_agent():
@wrap_model_call
async def before_model(request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse | AIMessage:
response = await handler(request)
await logger.ainfo("model node", response=response)
return response
model = FakeChatModel(responses=iter([
AIMessage(content="Test call tool", tool_calls=[
ToolCall(name="get_weather", args={"units": "celsius", "location": "TW", "include_forecast": True}, id="call_2"),
]),
AIMessage(content="Here you are!"),
]))
return create_agent(model=model, middleware=[before_model], tools=[get_weather, say_hello], checkpointer=None)
def get_tool_args(args):
return {k: v for k, v in args.items() if k!="runtime"}
async def _ahandle_stream(stream):
async for event in stream:
log=logger.bind(type=event["event"], name=event["name"])
data=event["data"]
if event["event"]=="on_chat_model_stream":
await log.ainfo("event data", data=event["data"])
elif event["event"]=="!on_chat_model_end":
await log.ainfo("event data", data=event["data"])
elif event["event"]=="on_tool_start":
if isinstance(data["input"], dict):
await log.ainfo("event data",
tool_name=event["name"],
input=get_tool_args(data["input"]),
tool_call_id=data["input"]["runtime"].tool_call_id,
data_keys=event["data"].keys(),
data_input_keys=data["input"].keys(),
)
else:
await log.ainfo("event data",
tool_name=event["name"],
input=data["input"],
data_keys=event["data"].keys(),
)
elif event["event"]=="on_tool_end":
await log.ainfo("event data",
tool_name=event["name"],
output=data["output"],
)
elif event["event"]=="on_tool_error":
await log.ainfo("event data",
tool_call_id=data["input"]["runtime"].tool_call_id,
input=get_tool_args(data["input"]),
error=event["data"]["error"],
data_keys=event["data"].keys(),
data_input_keys=data["input"].keys(),
)
else:
await log.ainfo("event data")
async def main():
log=logger
weather_agent = create_weather_agent()
supervisor = StateGraph(AgentState)
supervisor.add_node("weather_agent", weather_agent)
supervisor.add_edge(START, "weather_agent")
supervisor.add_edge("weather_agent", END)
agent = supervisor.compile(checkpointer=checkpoint)
config={"configurable": {"thread_id": uuid.uuid4().hex, "checkpoint_ns": ""}}
stream=agent.astream_events(
input={"messages": [HumanMessage(content="What's the weather in Tokyo? Include forecast")]},
config=config,
version="v2",
exclude_types=["chain"],
durability="exit",
)
await _ahandle_stream(stream)
state = await agent.aget_state(config)
await log.ainfo("state", state=state)
interrupts = state.interrupts if state.interrupts and len(state.interrupts) > 0 else []
c = None
async for c in agent.checkpointer.alist(config, limit=1):
await logger.ainfo("latest checkpoint", latest=c.config)
interrupts = state.interrupts if state.interrupts and len(state.interrupts) else []
def frontend_get_weather(kwargs):
units = kwargs["units"]
location = kwargs["location"]
include_forecast = kwargs["include_forecast"]
temp = 18 if units == "celsius" else 68
result = f"Current weather in {location}: {temp} degrees {units[0].upper()}"
if include_forecast:
result += "\nNext 5 days: Sunny"
return result
tool_message={
"id": "xyz",
"role": "tool",
"toolCallId": interrupts[0].value.tool_call_id,
"content": frontend_get_weather(interrupts[0].value.tool_call_args),
}
await logger.ainfo("resume")
stream=agent.astream_events(
input=Command(resume={
interrupts[0].id: ToolCallResume(
type=InterruptType.AGUI_TOOL_CALL,
tool_call_id=tool_message["toolCallId"],
content=tool_message["content"],
),
}),
config=c.config,
version="v2",
exclude_types=["chain"],
durability="exit",
)
await _ahandle_stream(stream)
state = await agent.aget_state(config)
for msg in state.values["messages"]:
print(f"{msg.__class__.__name__}: {msg.content}")
return
if __name__ == "__main__":
import asyncio
asyncio.run(main())
I was trying to resume a interrupt from a specific checkpoint_id, but the result seems to be replay instead of resume.
The expect output should be like:
AIMessage: Test call tool
ToolMessage: Current weather in TW: 18 degrees C
Next 5 days: Sunny
AIMessage: Here you are!
HumanMessage: What's the weather in Tokyo? Include forecast
AIMessage: Here you are!
Seems that the second run for resume still run from the beginning of the graph, not interrupt trigger point.
I tried following:
Checked other resources
Related Issues / PRs
No response
Reproduction Steps / Example Code (Python)
Error Message and Stack Trace (if applicable)
Description
I was trying to resume a interrupt from a specific checkpoint_id, but the result seems to be replay instead of resume.
The expect output should be like:
But got:
Seems that the second run for resume still run from the beginning of the graph, not interrupt trigger point.
I tried following:
System Info
System Information
Package Information