init
This commit is contained in:
@@ -0,0 +1,76 @@
|
||||
from typing import Any, Dict, List, Optional
|
||||
from langchain_core.messages import BaseMessage, HumanMessage
|
||||
from langchain_openai import ChatOpenAI
|
||||
from langgraph.graph import StateGraph, END
|
||||
from config import Config
|
||||
|
||||
|
||||
class AgentState:
|
||||
"""State definition for the agent workflow"""
|
||||
messages: List[BaseMessage]
|
||||
current_step: str
|
||||
context: Dict[str, Any]
|
||||
|
||||
def __init__(self, messages: List[BaseMessage] = None, current_step: str = "start", context: Dict[str, Any] = None):
|
||||
self.messages = messages or []
|
||||
self.current_step = current_step
|
||||
self.context = context or {}
|
||||
|
||||
|
||||
class BaseAgent:
|
||||
"""Base agent class with common functionality"""
|
||||
|
||||
def __init__(self, model_name: str = Config.DEFAULT_MODEL):
|
||||
self.model = ChatOpenAI(
|
||||
model=model_name,
|
||||
api_key=Config.OPENAI_API_KEY,
|
||||
temperature=0.1,
|
||||
max_retries=Config.MAX_RETRIES,
|
||||
timeout=Config.TIMEOUT
|
||||
)
|
||||
self.graph = self._build_graph()
|
||||
|
||||
def _build_graph(self) -> StateGraph:
|
||||
"""Build the state graph for the agent"""
|
||||
workflow = StateGraph(AgentState)
|
||||
|
||||
# Add nodes and edges
|
||||
workflow.add_node("process_input", self._process_input)
|
||||
workflow.add_node("generate_response", self._generate_response)
|
||||
|
||||
# Define edges
|
||||
workflow.add_edge("process_input", "generate_response")
|
||||
workflow.add_edge("generate_response", END)
|
||||
|
||||
# Set entry point
|
||||
workflow.set_entry_point("process_input")
|
||||
|
||||
return workflow.compile()
|
||||
|
||||
def _process_input(self, state: AgentState) -> AgentState:
|
||||
"""Process user input"""
|
||||
# This is a base implementation - subclasses should override
|
||||
state.current_step = "processed"
|
||||
return state
|
||||
|
||||
def _generate_response(self, state: AgentState) -> AgentState:
|
||||
"""Generate response using the LLM"""
|
||||
if state.messages:
|
||||
response = self.model.invoke(state.messages)
|
||||
state.messages.append(response)
|
||||
return state
|
||||
|
||||
def run(self, user_input: str, **kwargs) -> Dict[str, Any]:
|
||||
"""Run the agent with user input"""
|
||||
initial_state = AgentState(
|
||||
messages=[HumanMessage(content=user_input)],
|
||||
context=kwargs
|
||||
)
|
||||
|
||||
result = self.graph.invoke(initial_state)
|
||||
|
||||
return {
|
||||
"messages": result.messages,
|
||||
"context": result.context,
|
||||
"final_step": result.current_step
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
from typing import Dict, Any, List
|
||||
from langchain_core.messages import BaseMessage, HumanMessage, AIMessage
|
||||
from langgraph.graph import StateGraph, END
|
||||
from .base_agent import BaseAgent, AgentState
|
||||
|
||||
|
||||
class ConversationAgent(BaseAgent):
|
||||
"""Agent for handling multi-turn conversations"""
|
||||
|
||||
def __init__(self, model_name: str = None):
|
||||
super().__init__(model_name)
|
||||
self.conversation_history: List[BaseMessage] = []
|
||||
|
||||
def _build_graph(self) -> StateGraph:
|
||||
"""Build conversation-specific graph"""
|
||||
workflow = StateGraph(AgentState)
|
||||
|
||||
# Add nodes
|
||||
workflow.add_node("analyze_intent", self._analyze_intent)
|
||||
workflow.add_node("generate_response", self._generate_response)
|
||||
workflow.add_node("update_context", self._update_context)
|
||||
|
||||
# Define edges
|
||||
workflow.add_edge("analyze_intent", "generate_response")
|
||||
workflow.add_edge("generate_response", "update_context")
|
||||
workflow.add_edge("update_context", END)
|
||||
|
||||
# Set entry point
|
||||
workflow.set_entry_point("analyze_intent")
|
||||
|
||||
return workflow.compile()
|
||||
|
||||
def _analyze_intent(self, state: AgentState) -> AgentState:
|
||||
"""Analyze user intent and conversation context"""
|
||||
# Simple intent analysis - can be enhanced with more sophisticated logic
|
||||
user_message = state.messages[-1] if state.messages else None
|
||||
|
||||
if user_message and isinstance(user_message, HumanMessage):
|
||||
content = user_message.content.lower()
|
||||
|
||||
# Basic intent detection
|
||||
if any(word in content for word in ["hello", "hi", "hey", "greetings"]):
|
||||
state.context["intent"] = "greeting"
|
||||
elif any(word in content for word in ["help", "assist", "support"]):
|
||||
state.context["intent"] = "help"
|
||||
elif "?" in content:
|
||||
state.context["intent"] = "question"
|
||||
else:
|
||||
state.context["intent"] = "general"
|
||||
|
||||
state.current_step = "intent_analyzed"
|
||||
return state
|
||||
|
||||
def _generate_response(self, state: AgentState) -> AgentState:
|
||||
"""Generate response considering conversation history"""
|
||||
# Combine conversation history with current message
|
||||
all_messages = self.conversation_history + state.messages
|
||||
|
||||
if all_messages:
|
||||
response = self.model.invoke(all_messages)
|
||||
state.messages.append(response)
|
||||
|
||||
state.current_step = "response_generated"
|
||||
return state
|
||||
|
||||
def _update_context(self, state: AgentState) -> AgentState:
|
||||
"""Update conversation context and history"""
|
||||
# Add the conversation to history (excluding system messages)
|
||||
for message in state.messages:
|
||||
if isinstance(message, (HumanMessage, AIMessage)):
|
||||
self.conversation_history.append(message)
|
||||
|
||||
# Limit conversation history to avoid token limits
|
||||
if len(self.conversation_history) > 10:
|
||||
self.conversation_history = self.conversation_history[-10:]
|
||||
|
||||
state.current_step = "context_updated"
|
||||
return state
|
||||
|
||||
def run(self, user_input: str, **kwargs) -> Dict[str, Any]:
|
||||
"""Run conversation with history management"""
|
||||
initial_state = AgentState(
|
||||
messages=[HumanMessage(content=user_input)],
|
||||
context=kwargs
|
||||
)
|
||||
|
||||
result = self.graph.invoke(initial_state)
|
||||
|
||||
return {
|
||||
"messages": result.messages,
|
||||
"context": result.context,
|
||||
"conversation_history": self.conversation_history,
|
||||
"final_step": result.current_step
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
from typing import Dict, Any, List, Optional
|
||||
from langchain_core.messages import BaseMessage, HumanMessage, AIMessage, ToolMessage
|
||||
from langchain_core.tools import BaseTool
|
||||
from langgraph.graph import StateGraph, END
|
||||
from langgraph.prebuilt import ToolNode
|
||||
from .base_agent import BaseAgent, AgentState
|
||||
from tools.calculator import CalculatorTool
|
||||
from tools.web_search import WebSearchTool
|
||||
|
||||
|
||||
class ToolAgent(BaseAgent):
|
||||
"""Agent that can use tools to accomplish tasks"""
|
||||
|
||||
def __init__(self, model_name: str = None, tools: List[BaseTool] = None):
|
||||
# Initialize with default tools if none provided
|
||||
if tools is None:
|
||||
tools = [CalculatorTool(), WebSearchTool()]
|
||||
|
||||
self.tools = tools
|
||||
self.tool_node = ToolNode(tools)
|
||||
super().__init__(model_name)
|
||||
|
||||
def _build_graph(self) -> StateGraph:
|
||||
"""Build tool-using graph"""
|
||||
workflow = StateGraph(AgentState)
|
||||
|
||||
# Add nodes
|
||||
workflow.add_node("agent", self._agent_node)
|
||||
workflow.add_node("tools", self.tool_node)
|
||||
|
||||
# Define edges
|
||||
workflow.add_edge("tools", "agent")
|
||||
|
||||
# Conditional routing
|
||||
workflow.add_conditional_edges(
|
||||
"agent",
|
||||
self._should_use_tools,
|
||||
{
|
||||
"tools": "tools",
|
||||
"end": END,
|
||||
}
|
||||
)
|
||||
|
||||
# Set entry point
|
||||
workflow.set_entry_point("agent")
|
||||
|
||||
return workflow.compile()
|
||||
|
||||
def _agent_node(self, state: AgentState) -> AgentState:
|
||||
"""Agent node that decides whether to use tools"""
|
||||
# Bind tools to the model
|
||||
model_with_tools = self.model.bind_tools(self.tools)
|
||||
|
||||
# Get the last message
|
||||
if state.messages:
|
||||
response = model_with_tools.invoke(state.messages)
|
||||
state.messages.append(response)
|
||||
|
||||
return state
|
||||
|
||||
def _should_use_tools(self, state: AgentState) -> str:
|
||||
"""Determine if tools should be used"""
|
||||
last_message = state.messages[-1]
|
||||
|
||||
# If the last message has tool calls, route to tools
|
||||
if hasattr(last_message, 'tool_calls') and last_message.tool_calls:
|
||||
return "tools"
|
||||
|
||||
# Otherwise, end the workflow
|
||||
return "end"
|
||||
|
||||
def run(self, user_input: str, **kwargs) -> Dict[str, Any]:
|
||||
"""Run the tool-using agent"""
|
||||
initial_state = AgentState(
|
||||
messages=[HumanMessage(content=user_input)],
|
||||
context=kwargs
|
||||
)
|
||||
|
||||
result = self.graph.invoke(initial_state)
|
||||
|
||||
return {
|
||||
"messages": result.messages,
|
||||
"context": result.context,
|
||||
"tools_used": [tool.name for tool in self.tools],
|
||||
"final_step": result.current_step
|
||||
}
|
||||
Reference in New Issue
Block a user