LangGraph — Conditional Workflow with Decision Routing

 


LangState and StateGraph are different, but directly linked. This is actually one of the most important concepts in LangGraph.

LangState                  StateGraph

   │                                 │

   │ defines                    │ uses

   ▼                               ▼

"What data does          "How does the

the workflow carry?"      workflow run?"



class LangState(TypedDict):
    messages: Annotated[list, add_messages]


LangState = the data/state structure , "Every state in my graph has a messages field."


StateGraph = the workflow that operates on that state

builder = StateGraph(LangState)

There is the link between them.

we a're telling LangGraph:

"Build me a graph whose state follows LangState."




Comlpete code 


import json
from langgraph.graph import StateGraph, START, END
from typing import Annotated, TypedDict

from dotenv import load_dotenv
from langchain_core.messages import HumanMessage, ToolMessage
from langchain_core.tools import tool
from langchain_google_genai import ChatGoogleGenerativeAI

from langgraph.graph.message import add_messages


# Load environment variables
load_dotenv()


# -------------------------------------------------------------
# 1. Define Custom Python Tools
# -------------------------------------------------------------

@tool
def calculate_compound_interest(
    principal: float,
    annual_rate: float,
    years: int
) -> str:
    """Calculates compound interest."""

    amount = principal * ((1 + (annual_rate / 100)) ** years)
    interest = amount - principal

    return json.dumps({
        "principal": principal,
        "interest_earned": round(interest, 2),
        "total_amount": round(amount, 2),
        "years": years,
    })


@tool
def get_stock_price(ticker: str) -> str:
    """Fetches the mock current stock price."""

    mock_prices = {
        "AAPL": 220.50,
        "GOOGL": 175.30,
        "MSFT": 415.00
    }

    price = mock_prices.get(ticker.upper(), 100.00)

    return json.dumps({
        "ticker": ticker.upper(),
        "price_usd": price
    })


# -------------------------------------------------------------
# 2. Tools
# -------------------------------------------------------------

tools_list = [
    calculate_compound_interest,
    get_stock_price
]

tools_by_name = {
    t.name: t for t in tools_list
}


# -------------------------------------------------------------
# 3. Initialize Gemini
# -------------------------------------------------------------

llm = ChatGoogleGenerativeAI(
    model="gemini-2.5-flash",
    temperature=0.0,
)

llm_with_tools = llm.bind_tools(tools_list)


# =============================================================
# 4. LANGSTATE
# =============================================================

class LangState(TypedDict):
    messages: Annotated[list, add_messages]


# =============================================================
# 5. NODE 1 — LLM
# =============================================================

def call_llm(state: LangState):

    print("\n--- LLM NODE ---")

    ai_response = llm_with_tools.invoke(
        state["messages"]
    )

    print("LLM Response:", ai_response)

    return {
        "messages": [ai_response]
    }


# =============================================================
# 6. NODE 2 — TOOL
# =============================================================

def call_tools(state: LangState):

    print("\n--- TOOL NODE ---")

    last_message = state["messages"][-1]

    tool_messages = []

    for tool_call in last_message.tool_calls:

        tool_name = tool_call["name"]
        tool_args = tool_call["args"]
        tool_id = tool_call["id"]

        print("Tool:", tool_name)
        print("Arguments:", tool_args)

        selected_tool = tools_by_name[tool_name]

        tool_result = selected_tool.invoke(tool_args)

        print("Tool Result:", tool_result)

        tool_messages.append(
            ToolMessage(
                content=str(tool_result),
                tool_call_id=tool_id
            )
        )

    return {
        "messages": tool_messages
    }


# =============================================================
# 7. DECISION NODE / ROUTER
# =============================================================

def should_continue(state: LangState):

    print("\n--- DECISION ---")

    last_message = state["messages"][-1]

    # If LLM requested one or more tools,
    # route the workflow to the tool node.
    if last_message.tool_calls:

        print("Decision: Tool required")

        return "tools"

    # Otherwise, finish the workflow.
    print("Decision: No tool required")

    return "end"


# =============================================================
# 8. CREATE LANGGRAPH
# =============================================================

builder = StateGraph(LangState)


# Add nodes
builder.add_node("llm", call_llm)
builder.add_node("tools", call_tools)


# -------------------------------------------------------------
# Graph starts with the LLM
# -------------------------------------------------------------

builder.add_edge(START, "llm")


# -------------------------------------------------------------
# Conditional routing
# -------------------------------------------------------------

builder.add_conditional_edges(
    "llm",
    should_continue,
    {
        "tools": "tools",
        "end": END
    }
)


# -------------------------------------------------------------
# After executing a tool, go back to the LLM
# -------------------------------------------------------------

builder.add_edge("tools", "llm")


# Build the graph
graph = builder.compile()


# =============================================================
# 9. INITIAL LANGSTATE
# =============================================================

user_query = (
    "If I invest $5,000 at a 7% annual interest rate "
    " annual_rate is the annual percentage rate. "
 " For example, 7% annual rate should be passed as 7. "
    "for 10 years, how much will I earn? "
)

initial_state: LangState = {
    "messages": [
        HumanMessage(content=user_query)
    ]
}


# =============================================================
# 10. RUN LANGGRAPH
# =============================================================

print("\nInitial State:")
print(initial_state)

final_state = graph.invoke(initial_state)

print("\nFinal State:")
print(final_state)

Comments

Popular posts from this blog

Aggregate function with spring data

Java Persistence API with Spring Data

Thread , Runnable and ExecutorService