Spaces:
Sleeping
Sleeping
File size: 2,268 Bytes
b0540b3 |
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 |
from langgraph.graph import StateGraph
from src.state.state import State
from src.basic_bot.chatbot_node import ChatbotNode
from src.basic_bot.chatbot_with_tool_node import ChatbotWithToolNode
from langgraph.graph import StateGraph
from langgraph.prebuilt import tools_condition
from src.tools.search_tool import get_tools, create_tool_node
from src.basic_bot.chatbot_with_tool_node import ChatbotWithToolNode
class GraphBuilder:
"""
Manages the creation and setup of the StateGraph based on use cases.
"""
def __init__(self,model):
self.llm = model
self.graph_builder = StateGraph(State)
self.chatbot_node = ChatbotNode(model)
self.chatbot_with_tool_node = ChatbotWithToolNode(model)
def build_graph(self):
"""
Builds and returns the LangGraph graph based on the defined nodes and edges.
"""
# Initialize state graph
graph_builder = StateGraph(State)
# Define tools and tool node
tools = get_tools()
tool_node = create_tool_node(tools)
# Define LLM
llm = self.llm
# Define chatbot node
obj_chatbot_with_node = ChatbotWithToolNode(llm)
chatbot_node = obj_chatbot_with_node.create_chatbot(tools)
# Add nodes
graph_builder.add_node("chatbot", chatbot_node)
graph_builder.add_node("tools", tool_node)
# Define conditional and direct edges
graph_builder.add_conditional_edges("chatbot", tools_condition)
graph_builder.add_edge("tools", "chatbot")
# Set entry point and compile graph
graph_builder.set_entry_point("chatbot")
return graph_builder
def setup_graph(self, usecase: str):
"""
Sets up the graph for the selected use case.
"""
if usecase == "Basic Chatbot":
self.graph_builder.add_node("chatbot", self.chatbot_node.process)
self.graph_builder.set_entry_point("chatbot")
self.graph_builder.set_finish_point("chatbot")
elif usecase == "Chatbot with Tool":
self.graph_builder = self.build_graph()
else:
raise ValueError("Invalid use case selected.")
return self.graph_builder.compile()
|