"""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: - 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"))