rsobieski's picture
Update agent.py
57fcac6 verified
Raw
History Blame
3.86 kB
import os
import pandas as pd
from langchain_core.messages import HumanMessage, AIMessage
from langgraph.graph import StateGraph, MessagesState
from langchain_huggingface import HuggingFaceEndpoint
from tools import TOOLS
# --- Read local QA data 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()
}
def build_graph():
# Initialize Mistral model
llm = HuggingFaceEndpoint(
repo_id="mistralai/Mistral-7B-Instruct-v0.3",
task="conversational",
huggingfacehub_api_token=os.environ["HF_TOKEN"]
)
# Retriever node
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"]}
# Assistant node (LLM)
def assistant_node(state: MessagesState):
query = state["messages"][-1].content.strip()
system_prompt = (
"You are a helpful assistant evaluated by the GAIA benchmark.\n"
"Only return the final answer, with no explanations.\n"
"- No prefixes like 'Final answer:'\n"
"- If it's a list, output comma-separated\n"
"- If unknown, say 'Unknown'\n"
"- Never justify or explain"
)
# prompt = [
# {"role": "system", "content": system_prompt},
# {"role": "user", "content": query},
# ]
prompt = f"{system_prompt}\n\nUser: {query}\nAssistant:"
response = llm.invoke(prompt).strip()
# response = llm.invoke(messages)
for tag in ("Final answer:", "Answer:", "assistant:"):
if response.lower().startswith(tag.lower()):
response = response[len(tag):].strip()
return {"messages": [AIMessage(content=response.strip())]}
# Tool node
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()
# Agent class
class BasicAgent:
def __init__(self):
print("βœ… BasicAgent initialized with retriever + LLM + tools")
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