Final_Assignment / agent.py
shinichi123's picture
Upload 2 files
7faf98d verified
Raw History Blame Contribute Delete
6.84 kB
"""GAIA final assignment agent built with LangGraph.
LLM: Hugging Face Inference Providers (OpenAI-compatible router), paid from HF credits.
Required secret: HF_TOKEN (token with "Make calls to Inference Providers" permission)
Optional variable: MODEL_ID (default below)
"""
import contextlib
import io
import os
import re
import tempfile
import traceback
from typing import Annotated, TypedDict
import requests
from langchain_core.messages import AnyMessage, HumanMessage, SystemMessage
from langchain_core.tools import tool
from langchain_openai import ChatOpenAI
from langgraph.graph import START, StateGraph
from langgraph.graph.message import add_messages
from langgraph.prebuilt import ToolNode, tools_condition
API_URL = "https://agents-course-unit4-scoring.hf.space"
MODEL_ID = os.getenv("MODEL_ID", "openai/gpt-oss-120b")
SYSTEM_PROMPT = """You are a careful research assistant answering questions from the GAIA benchmark.
How to work:
- Think step by step. Use tools to look things up instead of guessing.
- Use web_search to find sources, then fetch_webpage to read the most promising page (Wikipedia is often useful).
- Use run_python for any calculation, counting, sorting, or for reading attached files (pandas is available).
- Read the question very carefully, including tricky details and the exact answer format requested.
How to answer:
- When you are done, end your reply with a line of the form:
FINAL ANSWER: <answer>
- <answer> must be as short as possible: a number, a few words, or a comma-separated list.
- Numbers: no thousands separators and no units unless the question asks for them.
- Strings: no articles or abbreviations unless asked; write digits as numerals unless told otherwise.
- Follow any formatting the question requests (ordering, capitalization, decimal places, etc.).
"""
# ---------------------------------------------------------------- tools
@tool
def web_search(query: str) -> str:
"""Search the web. Returns titles, URLs and short snippets of the top results."""
from ddgs import DDGS
try:
results = DDGS().text(query, max_results=6)
except Exception as e:
return f"Search error: {e}"
if not results:
return "No results."
return "\n\n".join(
f"{r.get('title')}\n{r.get('href')}\n{r.get('body')}" for r in results
)
@tool
def fetch_webpage(url: str, start: int = 0) -> str:
"""Fetch a web page and return its text as Markdown (tables included).
Output is limited to 15000 characters; pass start=15000, 30000, ... to read further."""
from bs4 import BeautifulSoup
from markdownify import markdownify
try:
resp = requests.get(url, timeout=20, headers={"User-Agent": "Mozilla/5.0"})
resp.raise_for_status()
except Exception as e:
return f"Fetch error: {e}"
soup = BeautifulSoup(resp.text, "html.parser")
for tag in soup(["script", "style", "nav", "footer", "header"]):
tag.decompose()
text = re.sub(r"\n{3,}", "\n\n", markdownify(str(soup)))
chunk = text[start : start + 15000]
if start + 15000 < len(text):
chunk += f"\n\n[... truncated, total {len(text)} chars. Use start={start + 15000} to continue]"
return chunk
_PY_GLOBALS: dict = {}
@tool
def run_python(code: str) -> str:
"""Execute Python code and return what it prints. Always use print() to show results.
pandas and openpyxl are available. Variables persist between calls within one question."""
buf = io.StringIO()
try:
with contextlib.redirect_stdout(buf):
exec(code, _PY_GLOBALS)
except Exception:
return buf.getvalue() + "\n" + traceback.format_exc(limit=2)
out = buf.getvalue()
return out[:10000] if out else "(no output - use print())"
TOOLS = [web_search, fetch_webpage, run_python]
# ---------------------------------------------------------------- graph
class State(TypedDict):
messages: Annotated[list[AnyMessage], add_messages]
def build_graph():
llm = ChatOpenAI(
model=MODEL_ID,
base_url="https://router.huggingface.co/v1",
api_key=os.environ["HF_TOKEN"],
temperature=0,
)
llm_with_tools = llm.bind_tools(TOOLS)
def assistant(state: State):
return {"messages": [llm_with_tools.invoke(state["messages"])]}
graph = StateGraph(State)
graph.add_node("assistant", assistant)
graph.add_node("tools", ToolNode(TOOLS))
graph.add_edge(START, "assistant")
graph.add_conditional_edges("assistant", tools_condition) # tool call -> tools, else END
graph.add_edge("tools", "assistant")
return graph.compile()
# ---------------------------------------------------------------- helpers
def download_task_file(task_id: str, file_name: str) -> str:
resp = requests.get(f"{API_URL}/files/{task_id}", timeout=30)
resp.raise_for_status()
path = os.path.join(tempfile.gettempdir(), file_name)
with open(path, "wb") as f:
f.write(resp.content)
return path
def clean_answer(content) -> str:
if isinstance(content, list): # some providers return content blocks
content = "".join(b.get("text", "") if isinstance(b, dict) else str(b) for b in content)
text = (content or "").strip()
match = re.search(r"FINAL ANSWER:\s*(.*)", text, re.S | re.I)
if match:
text = match.group(1).strip()
return text.strip().strip('"').rstrip(".").strip()
# ---------------------------------------------------------------- agent
class BasicAgent:
def __init__(self):
self.graph = build_graph()
print(f"BasicAgent initialized with model {MODEL_ID}")
def __call__(self, question: str, task_id: str | None = None, file_name: str | None = None) -> str:
_PY_GLOBALS.clear()
prompt = question
if task_id and file_name:
try:
path = download_task_file(task_id, file_name)
prompt += f"\n\nAttached file saved at: {path}\n(Read it with run_python.)"
except Exception as e:
prompt += f"\n\n(The attached file could not be downloaded: {e})"
try:
result = self.graph.invoke(
{"messages": [SystemMessage(SYSTEM_PROMPT), HumanMessage(prompt)]},
config={"recursion_limit": 30},
)
answer = clean_answer(result["messages"][-1].content)
except Exception as e:
print(f"Agent error: {e}")
answer = ""
print(f"Answer: {answer!r}")
return answer
# ---------------------------------------------------------------- local test
if __name__ == "__main__":
q = requests.get(f"{API_URL}/random-question", timeout=30).json()
print("Q:", q["question"])
print("file:", q.get("file_name"))
BasicAgent()(q["question"], q["task_id"], q.get("file_name"))