gaia-final-assignment / gaia_client.py
ykumar2020's picture
Publish verified modular GAIA agent source
c641d5f verified
Raw
History Blame Contribute Delete
10.9 kB
"""Typed client for the official Agents Course GAIA API contract."""
from __future__ import annotations
import hashlib
import json
from pathlib import Path
from typing import Any
import requests
from requests.adapters import HTTPAdapter
from urllib3.util.retry import Retry
from config import Settings
class GaiaClientError(RuntimeError):
pass
class GaiaNetworkError(GaiaClientError):
pass
class GaiaHTTPError(GaiaClientError):
pass
class GaiaPayloadError(GaiaClientError):
pass
class GaiaClient:
def __init__(self, settings: Settings, session: requests.Session | None = None):
self.settings = settings
self.session = session or requests.Session()
if hasattr(self.session, "headers"):
self.session.headers.setdefault("User-Agent", settings.user_agent)
retry = Retry(
total=settings.retries,
connect=settings.retries,
read=settings.retries,
status=settings.retries,
backoff_factor=settings.backoff_seconds,
status_forcelist=(429, 500, 502, 503, 504),
allowed_methods=frozenset({"GET"}),
raise_on_status=False,
)
adapter = HTTPAdapter(max_retries=retry)
self.session.mount("https://", adapter)
self.session.mount("http://", adapter)
def _get_json(self, path: str) -> Any:
try:
response = self.session.get(
f"{self.settings.api_url}{path}", timeout=self.settings.request_timeout
)
response.raise_for_status()
return response.json()
except requests.HTTPError as exc:
raise GaiaHTTPError(
f"GET {path} returned HTTP {exc.response.status_code}"
) from exc
except (requests.ConnectionError, requests.Timeout) as exc:
raise GaiaNetworkError(f"GET {path} failed: {exc}") from exc
except requests.JSONDecodeError as exc:
raise GaiaPayloadError(f"GET {path} returned invalid JSON") from exc
except requests.RequestException as exc:
raise GaiaNetworkError(f"GET {path} failed: {exc}") from exc
def get_questions(self) -> list[dict[str, Any]]:
data = self._get_json("/questions")
if not isinstance(data, list) or not data:
raise GaiaPayloadError(
"Questions endpoint returned an empty or invalid payload"
)
for item in data:
if (
not isinstance(item, dict)
or not item.get("task_id")
or item.get("question") is None
):
raise GaiaPayloadError("Questions endpoint returned a malformed task")
return data
def fetch_questions(self) -> list[dict[str, Any]]:
"""Backward-compatible alias."""
return self.get_questions()
def get_random_question(self) -> dict[str, Any]:
data = self._get_json("/random-question")
if (
not isinstance(data, dict)
or not data.get("task_id")
or data.get("question") is None
):
raise GaiaPayloadError("Random-question endpoint returned a malformed task")
return data
def download_task_file(self, task_id: str, filename: str) -> Path:
safe_name = Path(filename).name
if not safe_name:
raise ValueError("filename is required")
directory = self.settings.cache_dir / "attachments" / str(task_id)
destination = directory / safe_name
metadata_path = directory / "metadata.json"
directory.mkdir(parents=True, exist_ok=True)
def record_failure(error: str, content_type: str = "") -> None:
metadata_path.write_text(
json.dumps(
{
"task_id": str(task_id),
"filename": safe_name,
"content_type": content_type,
"status": "failed",
"error": error,
"source_url": f"{self.settings.api_url}/files/{task_id}",
},
indent=2,
),
encoding="utf-8",
)
if destination.is_file() and metadata_path.is_file():
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
checksum = hashlib.sha256(destination.read_bytes()).hexdigest()
if (
metadata.get("status") == "complete"
and metadata.get("sha256") == checksum
):
return destination
source_url = f"{self.settings.api_url}/files/{task_id}"
try:
response = self.session.get(
source_url,
timeout=max(self.settings.request_timeout, 120),
)
response.raise_for_status()
except requests.HTTPError as exc:
if not self.settings.hf_token:
record_failure(
f"HTTP {exc.response.status_code}; no HF_TOKEN for dataset fallback"
)
raise GaiaHTTPError(
f"Attachment {task_id} returned HTTP {exc.response.status_code}; "
"HF_TOKEN is required for the official GAIA dataset fallback"
) from exc
source_url = (
"https://huggingface.co/datasets/gaia-benchmark/GAIA/resolve/main/"
f"2023/validation/{safe_name}"
)
try:
response = self.session.get(
source_url,
timeout=max(self.settings.request_timeout, 120),
headers={"Authorization": f"Bearer {self.settings.hf_token}"},
)
response.raise_for_status()
except requests.RequestException as fallback_exc:
status = getattr(
getattr(fallback_exc, "response", None),
"status_code",
"network error",
)
record_failure(
f"Course endpoint HTTP {exc.response.status_code}; dataset fallback {status}"
)
access_hint = (
" Accept the GAIA dataset access conditions at "
"https://huggingface.co/datasets/gaia-benchmark/GAIA and use "
"a read-enabled HF_TOKEN."
if status in (401, 403)
else ""
)
raise GaiaHTTPError(
f"Attachment {task_id} unavailable from course endpoint and "
f"GAIA dataset fallback.{access_hint}"
) from fallback_exc
except (
requests.ConnectionError,
requests.Timeout,
requests.RequestException,
) as exc:
if not self.settings.hf_token:
record_failure(
f"{type(exc).__name__}: {exc}; no HF_TOKEN for dataset fallback"
)
raise GaiaNetworkError(f"Attachment {task_id} failed: {exc}") from exc
source_url = (
"https://huggingface.co/datasets/gaia-benchmark/GAIA/resolve/main/"
f"2023/validation/{safe_name}"
)
try:
response = self.session.get(
source_url,
timeout=max(self.settings.request_timeout, 120),
headers={"Authorization": f"Bearer {self.settings.hf_token}"},
)
response.raise_for_status()
except requests.RequestException as fallback_exc:
record_failure(
f"Course endpoint {type(exc).__name__}; dataset fallback "
f"{type(fallback_exc).__name__}"
)
raise GaiaNetworkError(
f"Attachment {task_id} failed from course endpoint and GAIA dataset fallback"
) from fallback_exc
content_type = response.headers.get("content-type", "application/octet-stream")
if "application/json" in content_type.lower():
record_failure("Attachment endpoint returned JSON", content_type)
raise GaiaPayloadError(
f"Attachment endpoint returned JSON: {response.text[:500]}"
)
if not response.content:
record_failure("Attachment was empty", content_type)
raise GaiaPayloadError(f"Attachment for {task_id} was empty")
temporary = destination.with_suffix(destination.suffix + ".tmp")
temporary.write_bytes(response.content)
temporary.replace(destination)
metadata = {
"task_id": str(task_id),
"filename": safe_name,
"content_type": content_type,
"size": len(response.content),
"sha256": hashlib.sha256(response.content).hexdigest(),
"status": "complete",
"source_url": source_url,
}
metadata_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8")
return destination
def download_attachment(self, task: dict[str, Any]) -> Path | None:
filename = str(task.get("file_name") or "").strip()
return (
self.download_task_file(str(task["task_id"]), filename)
if filename
else None
)
def submit_answers(
self, username: str, agent_code: str, answers: list[dict[str, str]]
) -> dict[str, Any]:
payload = {
"username": username.strip(),
"agent_code": agent_code,
"answers": answers,
}
try:
response = self.session.post(
f"{self.settings.api_url}/submit",
json=payload,
timeout=max(self.settings.request_timeout, 60),
)
response.raise_for_status()
data = response.json()
except requests.HTTPError as exc:
raise GaiaHTTPError(
f"Submission returned HTTP {exc.response.status_code}"
) from exc
except requests.JSONDecodeError as exc:
raise GaiaPayloadError("Submission returned invalid JSON") from exc
except requests.RequestException as exc:
raise GaiaNetworkError(f"Submission network failure: {exc}") from exc
if not isinstance(data, dict):
raise GaiaPayloadError("Submission endpoint returned an invalid payload")
return data
def submit(
self, username: str, agent_code: str, answers: list[dict[str, str]]
) -> dict[str, Any]:
"""Backward-compatible alias for the explicit submission action."""
return self.submit_answers(username, agent_code, answers)