3mpj commited on
Commit
f2845f9
·
verified ·
1 Parent(s): 5243e0d

Create agent.py

Browse files
Files changed (1) hide show
  1. agent.py +77 -0
agent.py ADDED
@@ -0,0 +1,77 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from dotenv import load_dotenv
3
+ from langgraph.graph import START, StateGraph, MessagesState
4
+ from langgraph.prebuilt import tools_condition
5
+ from langgraph.prebuilt import ToolNode
6
+ from langchain_google_genai import ChatGoogleGenerativeAI
7
+ from langchain_groq import ChatGroq
8
+ from langchain_huggingface import ChatHuggingFace, HuggingFaceEndpoint#, HuggingFaceEmbeddings
9
+ # from langchain_community.vectorstores import SupabaseVectorStore
10
+ from langchain_core.messages import SystemMessage, HumanMessage
11
+ # from langchain.tools.retriever import create_retriever_tool
12
+ # from supabase.client import Client, create_client
13
+
14
+ from prompt import SYSTEM_PROMPT
15
+ from tools import add, subtract, multiply, divide, web_search
16
+
17
+ load_dotenv()
18
+
19
+ HUGGINGFACEHUB_API_TOKEN = os.environ["HF_TOKEN"]
20
+
21
+ tools = [add, subtract, multiply, divide, web_search]
22
+
23
+
24
+ # Build graph function
25
+ def build_graph(provider: str = "huggingface") -> StateGraph:
26
+ """Build the graph"""
27
+ sys_msg = SystemMessage(content=SYSTEM_PROMPT)
28
+ if provider == "google":
29
+ # Google Gemini
30
+ llm = ChatGoogleGenerativeAI(model="gemini-2.0-flash", temperature=0)
31
+ elif provider == "groq":
32
+ # Groq https://console.groq.com/docs/models
33
+ llm = ChatGroq(model="qwen-qwq-32b", temperature=0) # optional : qwen-qwq-32b gemma2-9b-it
34
+ elif provider == "huggingface":
35
+ llm = ChatHuggingFace(
36
+ llm=HuggingFaceEndpoint(
37
+ repo_id="Qwen/Qwen2.5-Coder-32B-Instruct",
38
+ huggingfacehub_api_token=HUGGINGFACEHUB_API_TOKEN
39
+ ),
40
+ )
41
+ else:
42
+ raise ValueError("Invalid provider. Choose 'google', 'groq' or 'huggingface'.")
43
+ llm_with_tools = llm.bind_tools(tools)
44
+
45
+ # Node
46
+ def assistant(state: MessagesState):
47
+ """Assistant node"""
48
+ message = [sys_msg] + state["messages"]
49
+ return {"messages": [llm_with_tools.invoke(message)]}
50
+
51
+ builder = StateGraph(MessagesState)
52
+ builder.add_node("assistant", assistant)
53
+ builder.add_node("tools", ToolNode(tools))
54
+ builder.add_edge(START, "assistant")
55
+ builder.add_conditional_edges(
56
+ "assistant",
57
+ tools_condition,
58
+ )
59
+ builder.add_edge("tools", "assistant")
60
+
61
+ # Compile graph
62
+ return builder.compile()
63
+
64
+
65
+ class BasicAgent:
66
+ """A langgraph agent."""
67
+ def __init__(self):
68
+ print("BasicAgent initialized.")
69
+ self.graph = build_graph()
70
+
71
+ def __call__(self, question: str) -> str:
72
+ print(f"Agent received question (first 50 chars): {question[:50]}...")
73
+ # Wrap the question in a HumanMessage from langchain_core
74
+ messages = [HumanMessage(content=question)]
75
+ messages = self.graph.invoke({"messages": messages})
76
+ answer = messages['messages'][-1].content
77
+ return answer[14:]