Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/gaia/agents/base/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -4154,6 +4154,8 @@ def _python_to_json_type(py_type: str) -> str:
desc = param_info.get("description", "")
if desc:
prop["description"] = desc
if param_info.get("enum"):
prop["enum"] = list(param_info["enum"])
properties[param_name] = prop
if param_info.get("required", True):
required.append(param_name)
Expand Down
16 changes: 16 additions & 0 deletions src/gaia/agents/base/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,17 @@ def _infer_param_type(annotation: Any) -> str:
return "unknown"


def _literal_choices(annotation: Any) -> Optional[list]:
"""The values of a ``Literal[...]`` (or ``Optional[Literal[...]]``), else None."""
origin = typing.get_origin(annotation)
if origin is typing.Union or origin is types.UnionType:
non_none = [a for a in typing.get_args(annotation) if a is not types.NoneType]
return _literal_choices(non_none[0]) if len(non_none) == 1 else None
if origin is typing.Literal:
return list(typing.get_args(annotation))
return None


def _parse_arg_descriptions(docstring: Optional[str]) -> Dict[str, str]:
"""Extract per-argument text from a Google-style ``Args:`` block.

Expand Down Expand Up @@ -266,6 +277,11 @@ def decorator(f: Callable) -> Callable:
"type": _infer_param_type(annotation),
"required": param.default == inspect.Parameter.empty,
}
choices = _literal_choices(annotation)
if choices:
# An enum lets the server's tool-call grammar rule out bad values.
param_info["type"] = _infer_param_type(type(choices[0]))
param_info["enum"] = choices

description = arg_descriptions.get(name, "").strip()
if description:
Expand Down
3 changes: 2 additions & 1 deletion src/gaia/agents/tools/browser_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
import json
import logging
from pathlib import Path
from typing import Literal

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -46,7 +47,7 @@ def _ensure_web_client() -> bool:
@tool(atomic=True)
def fetch_page(
url: str,
extract: str = "text",
extract: Literal["text", "html", "links", "tables"] = "text",
max_length: int = 5000,
) -> str:
"""Fetch a web page and extract its content.
Expand Down
45 changes: 45 additions & 0 deletions tests/unit/test_tool_enum_args.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
"""A Literal-typed tool argument reaches the model as a JSON-schema enum."""

from types import SimpleNamespace
from typing import Literal, Optional

from gaia.agents.base.tools import tool


def test_literal_becomes_enum_in_the_registry():
registry = {}

@tool(registry=registry)
def pick(mode: Literal["text", "html"] = "text", other: Optional[str] = None):
"""Pick.

Args:
mode: Which.
other: Something else.
"""

params = registry["pick"]["parameters"]
assert params["mode"]["enum"] == ["text", "html"]
assert params["mode"]["type"] == "string"
assert "enum" not in params["other"]


def test_fetch_page_schema_lists_its_extract_modes():
from gaia.agents.base.agent import Agent
from gaia.agents.base.tools import _TOOL_REGISTRY
from gaia.agents.tools.browser_tools import BrowserToolsMixin

class Probe(BrowserToolsMixin):
_web_client = None

Probe().register_browser_tools()
agent = SimpleNamespace(
_tools_registry={"fetch_page": _TOOL_REGISTRY["fetch_page"]}
)
(schema,) = Agent._build_openai_tool_schemas(agent)
extract = schema["function"]["parameters"]["properties"]["extract"]
assert extract == {
"type": "string",
"description": extract["description"],
"enum": ["text", "html", "links", "tables"],
}
Loading