199 lines
4.4 KiB
Python
199 lines
4.4 KiB
Python
"""
|
|
LangGraph Example with ChatOpenAI.
|
|
|
|
A multi-turn conversational agent with memory and tool routing.
|
|
"""
|
|
|
|
from typing import TypedDict, List, Literal
|
|
from langgraph.graph import StateGraph, START, END
|
|
from langchain_openai import ChatOpenAI
|
|
from pydantic import BaseModel, Field
|
|
import os
|
|
from dotenv import load_dotenv
|
|
|
|
load_dotenv()
|
|
|
|
|
|
# ============ Model Helper ============
|
|
|
|
def create_chat_model(model: str = "openai/gpt-5.4", temperature: float = 0.7):
|
|
"""Create a ChatOpenAI model instance."""
|
|
return ChatOpenAI(
|
|
model=model,
|
|
temperature=temperature,
|
|
max_retries=2,
|
|
api_key=os.getenv("OPENAI_API_KEY"),
|
|
base_url=os.getenv("OPENAI_BASE_URL", "https://api.qnaigc.com/v1"),
|
|
)
|
|
|
|
|
|
# ============ State Schema ============
|
|
|
|
class Message(BaseModel):
|
|
role: str
|
|
content: str
|
|
|
|
|
|
class State(TypedDict):
|
|
messages: List[Message]
|
|
category: str
|
|
|
|
|
|
# ============ Router Function ============
|
|
|
|
def categorize_message(state: State) -> Literal["greeting", "question", "fallback"]:
|
|
"""Use LLM to categorize the user's message."""
|
|
|
|
model = create_chat_model(temperature=0.1)
|
|
|
|
last_message = state["messages"][-1].content if state["messages"] else ""
|
|
|
|
prompt = f"""
|
|
Categorize this message into ONE category: greeting, question, or fallback.
|
|
Return ONLY the category name.
|
|
|
|
Message: {last_message}
|
|
"""
|
|
|
|
response = model.invoke(prompt).content.strip().lower()
|
|
|
|
if "greeting" in response:
|
|
return "greeting"
|
|
elif "question" in response:
|
|
return "question"
|
|
else:
|
|
return "fallback"
|
|
|
|
|
|
# ============ Node Functions ============
|
|
|
|
def handle_greeting(state: State) -> dict:
|
|
"""Handle greeting messages."""
|
|
print(" [Node: Greeting]")
|
|
|
|
model = create_chat_model()
|
|
|
|
last_message = state["messages"][-1].content
|
|
|
|
prompt = f"""
|
|
Respond to this greeting in a friendly way. Keep it brief.
|
|
|
|
User: {last_message}
|
|
Assistant:
|
|
"""
|
|
|
|
response = model.invoke(prompt).content
|
|
|
|
return {
|
|
"messages": state["messages"] + [Message(role="assistant", content=response)]
|
|
}
|
|
|
|
|
|
def handle_question(state: State) -> dict:
|
|
"""Handle question messages."""
|
|
print(" [Node: Question]")
|
|
|
|
model = create_chat_model()
|
|
|
|
last_message = state["messages"][-1].content
|
|
|
|
prompt = f"""
|
|
Answer this question helpfully and concisely.
|
|
|
|
User: {last_message}
|
|
Assistant:
|
|
"""
|
|
|
|
response = model.invoke(prompt).content
|
|
|
|
return {
|
|
"messages": state["messages"] + [Message(role="assistant", content=response)]
|
|
}
|
|
|
|
|
|
def handle_fallback(state: State) -> dict:
|
|
"""Handle unrecognized messages."""
|
|
print(" [Node: Fallback]")
|
|
|
|
model = create_chat_model()
|
|
|
|
last_message = state["messages"][-1].content
|
|
|
|
prompt = f"""
|
|
Respond to this message in a helpful way.
|
|
|
|
User: {last_message}
|
|
Assistant:
|
|
"""
|
|
|
|
response = model.invoke(prompt).content
|
|
|
|
return {
|
|
"messages": state["messages"] + [Message(role="assistant", content=response)]
|
|
}
|
|
|
|
|
|
# ============ Build Graph ============
|
|
|
|
def create_chat_graph():
|
|
"""Create and compile the chat graph."""
|
|
|
|
builder = StateGraph(State)
|
|
|
|
# Add nodes
|
|
builder.add_node("greeting", handle_greeting)
|
|
builder.add_node("question", handle_question)
|
|
builder.add_node("fallback", handle_fallback)
|
|
|
|
# Add conditional routing
|
|
builder.add_conditional_edges(
|
|
START,
|
|
categorize_message,
|
|
{
|
|
"greeting": "greeting",
|
|
"question": "question",
|
|
"fallback": "fallback"
|
|
}
|
|
)
|
|
|
|
# All paths lead to END
|
|
builder.add_edge("greeting", END)
|
|
builder.add_edge("question", END)
|
|
builder.add_edge("fallback", END)
|
|
|
|
return builder.compile()
|
|
|
|
|
|
# ============ Main ============
|
|
|
|
def main():
|
|
"""Run the chat agent."""
|
|
graph = create_chat_graph()
|
|
|
|
print("=== LangGraph Chat Agent with ChatOpenAI ===\n")
|
|
|
|
# Test messages
|
|
test_inputs = [
|
|
"你好,很高兴见到你!",
|
|
"Python 中如何读取 JSON 文件?",
|
|
"随便聊聊吧"
|
|
]
|
|
|
|
for user_input in test_inputs:
|
|
print(f"User: {user_input}")
|
|
|
|
initial_state = {
|
|
"messages": [Message(role="user", content=user_input)],
|
|
"category": ""
|
|
}
|
|
|
|
result = graph.invoke(initial_state)
|
|
|
|
assistant_message = result["messages"][-1].content
|
|
print(f"Assistant: {assistant_message}\n")
|
|
print("-" * 50 + "\n")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|