update main.py
This commit is contained in:
@@ -1,7 +1,7 @@
|
|||||||
"""
|
"""
|
||||||
Main entry point for the agent with memory and human‑in‑the‑loop confirmation.
|
Main entry point for the agent with memory and human‑in‑the‑loop confirmation.
|
||||||
|
|
||||||
The agent is built on top of LangChain's `create_agent` API. It uses a
|
The agent is built on top of LangChain's `create_tool_calling_agent` API. It uses a
|
||||||
`MemorySaver` checkpoint to keep conversation history across calls, and it
|
`MemorySaver` checkpoint to keep conversation history across calls, and it
|
||||||
is configured with `interrupt_before=["tools"]` so that the agent pauses just
|
is configured with `interrupt_before=["tools"]` so that the agent pauses just
|
||||||
before invoking any tool. The pause allows us to ask the user for explicit
|
before invoking any tool. The pause allows us to ask the user for explicit
|
||||||
@@ -24,13 +24,13 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
import json
|
import json
|
||||||
from typing import Any, Dict, Iterable, Tuple
|
from typing import Any, Dict
|
||||||
|
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_core.messages import HumanMessage, SystemMessage
|
from langchain_core.messages import HumanMessage, SystemMessage
|
||||||
from langgraph.checkpoint.memory import MemorySaver
|
from langgraph.checkpoint.memory import MemorySaver
|
||||||
from langgraph.types import Command
|
from langgraph.types import Command
|
||||||
from langchain.agents import create_agent
|
from langchain.agents import create_tool_calling_agent
|
||||||
from langchain.tools import tool
|
from langchain.tools import tool
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
|
|
||||||
@@ -69,11 +69,11 @@ async def get_price(query: Dict[str, Any]) -> str:
|
|||||||
return f"The price of {product} is 42.00 {currency}."
|
return f"The price of {product} is 42.00 {currency}."
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Agent definition using create_agent (LangChain)
|
# Agent definition using create_tool_calling_agent (LangChain)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
memory = MemorySaver()
|
memory = MemorySaver()
|
||||||
|
|
||||||
agent = create_agent(
|
agent = create_tool_calling_agent(
|
||||||
llm=llm,
|
llm=llm,
|
||||||
tools=[get_price],
|
tools=[get_price],
|
||||||
system_prompt="You are a helpful assistant that can query prices.",
|
system_prompt="You are a helpful assistant that can query prices.",
|
||||||
@@ -81,9 +81,6 @@ agent = create_agent(
|
|||||||
interrupt_before=["tools"], # pause before any tool call
|
interrupt_before=["tools"], # pause before any tool call
|
||||||
)
|
)
|
||||||
|
|
||||||
# Compile the agent into a graph with checkpointing.
|
|
||||||
graph = agent.compile(checkpointer=memory, interrupt_before=["tools"])
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Helper to run the agent and pause before each tool call.
|
# Helper to run the agent and pause before each tool call.
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -98,10 +95,10 @@ async def ask_and_run(user_input: Dict[str, Any], config: Dict[str, Any]):
|
|||||||
Configuration dictionary that must contain ``configurable`` with
|
Configuration dictionary that must contain ``configurable`` with
|
||||||
``thread_id``.
|
``thread_id``.
|
||||||
"""
|
"""
|
||||||
async for chunk in graph.stream(user_input, config=config, stream_mode=["messages", "updates"]):
|
async for chunk in agent.stream(user_input, config=config, stream_mode=["messages", "updates"]):
|
||||||
# ``chunk`` is a tuple (type, data).
|
# ``chunk`` is a tuple (type, data).
|
||||||
chunk_type, chunk_data = chunk
|
chunk_type, chunk_data = chunk
|
||||||
state = graph.get_state(config)
|
state = agent.get_state(config)
|
||||||
|
|
||||||
if chunk_type == "messages":
|
if chunk_type == "messages":
|
||||||
# Stream token by token.
|
# Stream token by token.
|
||||||
@@ -120,7 +117,7 @@ async def ask_and_run(user_input: Dict[str, Any], config: Dict[str, Any]):
|
|||||||
console.print("[red]Action cancelled by user.[/]")
|
console.print("[red]Action cancelled by user.[/]")
|
||||||
break
|
break
|
||||||
# Resume from the same state.
|
# Resume from the same state.
|
||||||
await graph.ainvoke(Command(resume=None), config=config)
|
await agent.ainvoke(Command(resume=None), config=config)
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Main loop – three examples as requested.
|
# Main loop – three examples as requested.
|
||||||
|
|||||||
Reference in New Issue
Block a user