rsobieski's picture
Update agent.py
a72ec7c verified
Raw
History Blame
3.83 kB
import os
import pandas as pd
from langchain_core.messages import HumanMessage, AIMessage
from langgraph.graph import StateGraph, MessagesState
from langgraph.prebuilt import ToolNode
from langchain_huggingface import HuggingFaceEndpoint
from tools import TOOLS # your dictionary of tool functions
# --- Load local QA metadata for retriever ---
QA_PATH = "metadata.jsonl"
qa_pairs = pd.read_json(QA_PATH, lines=True)
qa_dict = {
row["Question"].strip(): row["Final answer"].strip()
for _, row in qa_pairs.iterrows()
}
# === LangGraph builder ===
def build_graph():
llm = HuggingFaceEndpoint(
repo_id="mistralai/Mistral-7B-Instruct-v0.3",
task="text-generation",
huggingfacehub_api_token=os.environ["HF_TOKEN"]
)
# Node 1: Retriever
def retriever_node(state: MessagesState):
query = state["messages"][-1].content.strip()
if query in qa_dict:
print("βœ… Exact match found in retriever.")
return {"messages": [AIMessage(content=qa_dict[query])]}
print("πŸ” No match. Sending to LLM.")
return {"messages": state["messages"]}
# Node 2: Assistant (LLM response parsing)
def assistant_node(state: MessagesState):
query = state["messages"][-1].content.strip()
prompt = (
"You are a helpful assistant for GAIA benchmark.\n"
"If you can answer directly, output ONLY the answer.\n"
"If you need to use a tool, reply in this format:\n"
"use_tool: <tool_name>: <tool_input>\n"
"Never explain anything."
)
full_prompt = f"{prompt}\n\nQuestion: {query}\nAnswer:"
response = llm.invoke(full_prompt).strip()
print(f"🧠 LLM said: {response}")
if response.startswith("use_tool:"):
return {
"messages": state["messages"] + [AIMessage(content=response)],
"tool_call": response # carry tool signal
}
else:
return {"messages": [AIMessage(content=response)]}
# Node 3: Tool execution
def tool_node(state: MessagesState):
try:
tool_signal = state.get("tool_call", "")
_, tool_name, tool_input = tool_signal.split(":", 2)
tool_name = tool_name.strip()
tool_input = tool_input.strip()
tool_fn = TOOLS.get(tool_name)
if not tool_fn:
print(f"❌ Unknown tool: {tool_name}")
return {"messages": [AIMessage(content="Unknown")]}
print(f"πŸ”§ Using tool: {tool_name} with input: {tool_input}")
tool_result = tool_fn(tool_input)
return {"messages": [AIMessage(content=str(tool_result))]}
except Exception as e:
print(f"⚠️ Tool error: {e}")
return {"messages": [AIMessage(content="Unknown")]} # fail-safe
# Build LangGraph
builder = StateGraph(MessagesState)
builder.add_node("retriever", retriever_node)
builder.add_node("assistant", assistant_node)
builder.add_node("tool", tool_node)
builder.set_entry_point("retriever")
builder.add_edge("retriever", "assistant")
builder.add_edge("assistant", "tool")
builder.add_edge("tool", "assistant")
builder.set_finish_point("assistant")
return builder.compile()
# === BasicAgent wrapper ===
class BasicAgent:
def __init__(self):
print("BasicAgent initialized with retriever + LLM + manual tool logic.")
self.graph = build_graph()
def __call__(self, question: str) -> str:
print(f"πŸ“₯ Question: {question[:100]}")
result = self.graph.invoke({"messages": [HumanMessage(content=question)]})
answer = result["messages"][-1].content.strip()
print(f"πŸ“€ Answer: {answer}")
return answer