Add interrupt info to graph repr

This commit is contained in:
Nuno Campos
2024-08-02 12:20:56 -07:00
parent 742f17689e
commit b1d0dbac77
3 changed files with 99 additions and 2 deletions
+7 -2
View File
@@ -493,6 +493,11 @@ class CompiledGraph(Pregel):
for key, n in self.builder.nodes.items():
node = n.runnable
metadata = n.metadata or {}
if key in self.interrupt_before_nodes:
metadata["__interrupt"] = "before"
elif key in self.interrupt_after_nodes:
metadata["__interrupt"] = "after"
if xray:
subgraph = (
node.get_graph(
@@ -509,11 +514,11 @@ class CompiledGraph(Pregel):
subgraph, prefix=key
)
else:
n = graph.add_node(node, key)
n = graph.add_node(node, key, metadata=metadata or None)
start_nodes[key] = n
end_nodes[key] = n
else:
n = graph.add_node(node, key, metadata=n.metadata)
n = graph.add_node(node, key, metadata=metadata or None)
start_nodes[key] = n
end_nodes[key] = n
for start, end in sorted(self.builder._all_edges):
@@ -508,6 +508,95 @@
'''
# ---
# name: test_conditional_graph.5
dict({
'edges': list([
dict({
'source': '__start__',
'target': 'agent',
}),
dict({
'source': 'tools',
'target': 'agent',
}),
dict({
'conditional': True,
'data': 'continue',
'source': 'agent',
'target': 'tools',
}),
dict({
'conditional': True,
'data': 'exit',
'source': 'agent',
'target': '__end__',
}),
]),
'nodes': list([
dict({
'data': '__start__',
'id': '__start__',
'type': 'schema',
}),
dict({
'data': dict({
'id': list([
'langchain',
'schema',
'runnable',
'RunnableAssign',
]),
'name': 'agent',
}),
'id': 'agent',
'metadata': dict({
'__interrupt': 'after',
}),
'type': 'runnable',
}),
dict({
'data': dict({
'id': list([
'langgraph',
'utils',
'RunnableCallable',
]),
'name': 'tools',
}),
'id': 'tools',
'metadata': dict({
'variant': 'b',
'version': 2,
}),
'type': 'runnable',
}),
dict({
'data': '__end__',
'id': '__end__',
'type': 'schema',
}),
]),
})
# ---
# name: test_conditional_graph.6
'''
%%{init: {'flowchart': {'curve': 'linear'}}}%%
graph TD;
__start__([__start__]):::first
agent(agent<hr/><small><em>__interrupt = after</em></small>)
tools(tools<hr/><small><em>version = 2
variant = b</em></small>)
__end__([__end__]):::last
__start__ --> agent;
tools --> agent;
agent -. &nbspcontinue&nbsp .-> tools;
agent -. &nbspexit&nbsp .-> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_conditional_state_graph
'{"title": "LangGraphInput", "type": "object", "properties": {"input": {"title": "Input", "type": "string"}, "agent_outcome": {"title": "Agent Outcome", "anyOf": [{"$ref": "#/definitions/AgentAction"}, {"$ref": "#/definitions/AgentFinish"}]}, "intermediate_steps": {"title": "Intermediate Steps", "type": "array", "items": {"type": "array", "minItems": 2, "maxItems": 2, "items": [{"$ref": "#/definitions/AgentAction"}, {"type": "string"}]}}}, "definitions": {"AgentAction": {"title": "AgentAction", "description": "Represents a request to execute an action by an agent.\\n\\nThe action consists of the name of the tool to execute and the input to pass\\nto the tool. The log is used to pass along extra information about the action.", "type": "object", "properties": {"tool": {"title": "Tool", "type": "string"}, "tool_input": {"title": "Tool Input", "anyOf": [{"type": "string"}, {"type": "object"}]}, "log": {"title": "Log", "type": "string"}, "type": {"title": "Type", "default": "AgentAction", "enum": ["AgentAction"], "type": "string"}}, "required": ["tool", "tool_input", "log"]}, "AgentFinish": {"title": "AgentFinish", "description": "Final return value of an ActionAgent.\\n\\nAgents return an AgentFinish when they have reached a stopping condition.", "type": "object", "properties": {"return_values": {"title": "Return Values", "type": "object"}, "log": {"title": "Log", "type": "string"}, "type": {"title": "Type", "default": "AgentFinish", "enum": ["AgentFinish"], "type": "string"}}, "required": ["return_values", "log"]}}}'
# ---
+3
View File
@@ -1931,6 +1931,9 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
)
config = {"configurable": {"thread_id": "1"}}
assert app_w_interrupt.get_graph().to_json() == snapshot
assert app_w_interrupt.get_graph().draw_mermaid() == snapshot
assert [
c for c in app_w_interrupt.stream({"input": "what is weather in sf"}, config)
] == [