rsobieski commited on
Commit
63f13b5
·
verified ·
1 Parent(s): 6d7f920

Update agent.py

Browse files
Files changed (1) hide show
  1. agent.py +61 -1
agent.py CHANGED
@@ -1 +1,61 @@
1
- from langgraph.graph import START, StateGraph, MessagesState
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """LangGraph agent using retriever fallback with LLM + tools."""
2
+ import os
3
+ from langgraph.graph import StateGraph, MessagesState
4
+ from langgraph.prebuilt import ToolNode, tools_condition
5
+ from langchain_core.messages import HumanMessage, AIMessage
6
+ from langchain_google_genai import ChatGoogleGenerativeAI
7
+ from langchain_core.runnables import Runnable
8
+ from tools import TOOLS
9
+ import pandas as pd
10
+
11
+ # Load metadata from local jsonl
12
+ QA_PATH = "/mnt/data/metadata.jsonl"
13
+ qa_pairs = pd.read_json(QA_PATH, lines=True)
14
+ qa_dict = {row["Question"].strip(): row["Final answer"].strip() for _, row in qa_pairs.iterrows()}
15
+
16
+ def build_graph():
17
+ """Construct a LangGraph agent with a QA retriever and fallback LLM+tools."""
18
+
19
+ # Initialize the LLM (e.g., Gemini Flash, zero temperature)
20
+ # llm = ChatGoogleGenerativeAI(model="gemini-1.5-flash", temperature=0)
21
+ llm = ChatGoogleGenerativeAI(model="gemini-2.5-flash", temperature=0)
22
+ llm_with_tools = llm.bind_tools(TOOLS)
23
+
24
+ # Step 1: Retriever node
25
+ def retriever_node(state: MessagesState):
26
+ query = state["messages"][-1].content.strip()
27
+ if query in qa_dict:
28
+ print(f"✅ Exact match found in retriever.")
29
+ return {"messages": [AIMessage(content=qa_dict[query])]}
30
+ print(f"🔍 No match found. Falling back to LLM.")
31
+ return {"messages": state["messages"]} # Continue to LLM if no match
32
+
33
+ # Step 2: LLM + Tools node
34
+ def assistant_node(state: MessagesState):
35
+ return {"messages": [llm_with_tools.invoke(state["messages"])]}
36
+
37
+ # Build LangGraph
38
+ builder = StateGraph(MessagesState)
39
+ builder.add_node("retriever", retriever_node)
40
+ builder.add_node("assistant", assistant_node)
41
+ builder.add_node("tools", ToolNode(TOOLS))
42
+
43
+ # Edges
44
+ builder.set_entry_point("retriever")
45
+ builder.add_edge("retriever", "assistant")
46
+ builder.add_conditional_edges("assistant", tools_condition)
47
+ builder.add_edge("tools", "assistant")
48
+ builder.set_finish_point("assistant")
49
+
50
+ return builder.compile()
51
+
52
+ # Final agent interface
53
+ class BasicAgent:
54
+ def __init__(self):
55
+ print("BasicAgent initialized with retriever + LLM.")
56
+ self.graph = build_graph()
57
+
58
+ def __call__(self, question: str) -> str:
59
+ print(f"Agent received question: {question[:80]}")
60
+ result = self.graph.invoke({"messages": [HumanMessage(content=question)]})
61
+ return result['messages'][-1].content.strip()