jstoppa/langgraph_basic_example
0
1from langgraph.graph import StateGraph2from typing import TypedDict, Annotated3from langgraph.graph.message import add_messages4from langchain_core.runnables.graph import MermaidDrawMethod5 6class State(TypedDict):7 messages: Annotated[list[str], add_messages]8 current_step: str9 10def collect_info(state: State) -> dict:11 print("\n--> In collect_info")12 print(f"Messages before: {state['messages']}")13 14 messages = state["messages"] + ["Information collected"]15 print(f"Messages after: {messages}")16 17 return {18 "messages": messages,19 "current_step": "process"20 }21 22def process_info(state: State) -> dict:23 print("\n--> In process_info")24 print(f"Messages before: {state['messages']}")25 26 messages = state["messages"] + ["Information processed"]27 print(f"Messages after: {messages}")28 29 return {30 "messages": messages,31 "current_step": "end"32 }33 34# Create and setup graph35workflow = StateGraph(State)36 37# Add nodes38workflow.add_node("collect", collect_info)39workflow.add_node("process", process_info)40 41# Add edges42workflow.add_edge("collect", "process")43 44# Set entry and finish points45workflow.set_entry_point("collect")46workflow.set_finish_point("process")47 48app = workflow.compile()49 50# Run workflow51print("\nStarting workflow...")52initial_state = State(messages=["Starting"], current_step="collect")53final_state = app.invoke(initial_state)54print(f"\nFinal messages: {final_state['messages']}")55 56# Save the graph visualization as PNG57png_data = app.get_graph().draw_mermaid_png(draw_method=MermaidDrawMethod.API)58with open("workflow_graph.png", "wb") as f:59 f.write(png_data)60print("\nGraph visualization saved as 'workflow_graph.png'")