from dataclasses import dataclass import os from pathlib import Path import shutil from typing import Any import requests from huggingface_hub import hf_hub_download @dataclass class TaskRecord: task_id: str question: str level: str file_name: str = "" class CourseAPIClient: def __init__(self, api_url: str): self.api_url = api_url.rstrip("/") def fetch_questions(self) -> list[TaskRecord]: response = requests.get(f"{self.api_url}/questions", timeout=30) response.raise_for_status() return [self._to_task_record(item) for item in response.json()] def fetch_random_question(self) -> TaskRecord: response = requests.get(f"{self.api_url}/random-question", timeout=30) response.raise_for_status() return self._to_task_record(response.json()) def fetch_attachment(self, task: TaskRecord, target_dir: Path) -> Path | None: if not task.file_name: return None target_dir.mkdir(parents=True, exist_ok=True) response = requests.get(f"{self.api_url}/files/{task.task_id}", timeout=60) output_path = target_dir / task.file_name if response.status_code == 200: output_path.write_bytes(response.content) return output_path return self._fetch_attachment_from_gaia_dataset(task, output_path) @staticmethod def _fetch_attachment_from_gaia_dataset( task: TaskRecord, output_path: Path, ) -> Path | None: candidate_paths = [ f"2023/validation/{task.file_name}", f"2023/test/{task.file_name}", ] for repo_file in candidate_paths: try: cached_path = hf_hub_download( repo_id="gaia-benchmark/GAIA", repo_type="dataset", filename=repo_file, token=os.getenv("HF_TOKEN") or None, ) except Exception: continue shutil.copyfile(cached_path, output_path) return output_path return None def submit_answers( self, username: str, agent_code: str, answers: list[dict[str, str]], ) -> dict[str, Any]: payload = { "username": username, "agent_code": agent_code, "answers": answers, } response = requests.post(f"{self.api_url}/submit", json=payload, timeout=120) response.raise_for_status() return response.json() @staticmethod def _to_task_record(item: dict[str, Any]) -> TaskRecord: return TaskRecord( task_id=item["task_id"], question=item["question"], level=str(item.get("Level", "")), file_name=item.get("file_name", "") or "", )