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
13 changes: 12 additions & 1 deletion src/any_agent/frameworks/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,7 +145,18 @@ async def get_response(
response_usage.prompt_tokens_details, "cached_tokens", 0
)
)
or 0
or 0,
# Required since openai 2.x. Older versions of the model allow
# extra fields, so passing it stays compatible with those.
cache_write_tokens=(
getattr(response_usage, "prompt_tokens_details", None)
and getattr(
response_usage.prompt_tokens_details,
"cache_write_tokens",
0,
)
)
or 0,
),
output_tokens_details=OutputTokensDetails(
reasoning_tokens=(
Expand Down
5 changes: 3 additions & 2 deletions src/any_agent/tools/wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@

if TYPE_CHECKING:
from agents import Tool as AgentTool
from google.adk.tools import BaseTool as GoogleTool
from google.adk.tools.base_tool import BaseTool as GoogleTool
from langchain_core.tools import BaseTool as LangchainTool
from llama_index.core.tools import FunctionTool as LlamaIndexTool
from smolagents import Tool as SmolagentsTool
Expand Down Expand Up @@ -87,7 +87,8 @@ def _wrap_tool_llama_index(tool: "Tool | LlamaIndexTool") -> "LlamaIndexTool":


def _wrap_tool_google(tool: "Tool | GoogleTool") -> "GoogleTool":
from google.adk.tools import BaseTool, FunctionTool
from google.adk.tools.base_tool import BaseTool
from google.adk.tools.function_tool import FunctionTool

if isinstance(tool, BaseTool):
return tool
Expand Down
2 changes: 1 addition & 1 deletion tests/integration/tools/test_wrap_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
import pytest
from agents.tool import Tool as OpenaiClass
from any_llm.utils.aio import run_async_in_sync
from google.adk.tools import FunctionTool as GoogleClass
from google.adk.tools.function_tool import FunctionTool as GoogleClass
from langchain_core.tools import BaseTool as LangchainClass
from llama_index.core.tools import FunctionTool as LlamaindexClass
from smolagents.tools import Tool as SmolagentsClass
Expand Down
6 changes: 3 additions & 3 deletions tests/unit/frameworks/test_google.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@


def test_load_google_default() -> None:
from google.adk.tools import FunctionTool
from google.adk.tools.function_tool import FunctionTool

mock_agent = MagicMock()
mock_model = MagicMock()
Expand All @@ -20,7 +20,7 @@ def __new__(cls, *args: Any, **kwargs: Any) -> MagicMock:
with (
patch("any_agent.frameworks.google.LlmAgent", mock_agent),
patch("any_agent.frameworks.google.DEFAULT_MODEL_TYPE", mock_model),
patch("google.adk.tools.FunctionTool", MockedFunctionTool),
patch("google.adk.tools.function_tool.FunctionTool", MockedFunctionTool),
):
AnyAgent.create(
AgentFramework.GOOGLE, AgentConfig(model_id="mistral:mistral-small-latest")
Expand Down Expand Up @@ -65,7 +65,7 @@ def test_run_google_custom_args() -> None:
patch("any_agent.frameworks.google.LlmAgent", mock_agent),
patch("any_agent.frameworks.google.InMemoryRunner", mock_runner),
patch("any_agent.frameworks.google.DEFAULT_MODEL_TYPE"),
patch("google.adk.tools.FunctionTool"),
patch("google.adk.tools.function_tool.FunctionTool"),
):
agent = AnyAgent.create(
AgentFramework.GOOGLE, AgentConfig(model_id="mistral:mistral-small-latest")
Expand Down
5 changes: 4 additions & 1 deletion tests/unit/frameworks/test_smolagents.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@

import pytest

from any_llm import LLMProvider

from any_agent import AgentConfig, AgentFramework, AnyAgent
from any_agent.frameworks.smolagents import AnyLLMModel

Expand Down Expand Up @@ -131,7 +133,8 @@ def test_parses_model_id_correctly(self) -> None:
api_base="https://api.example.com",
)

assert model._provider.value == "openai"
# any-llm 1.28 widened split_model_provider to return `str | LLMProvider`.
assert LLMProvider(model._provider) is LLMProvider.OPENAI
assert model._anyllm_completion_kwargs["model"] == "gpt-4o"
assert model._api_key == "test-key"
assert model._api_base == "https://api.example.com"
Expand Down
12 changes: 7 additions & 5 deletions tests/unit/frameworks/test_tinyagent.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,18 +250,20 @@ def create_mock_response(content: str, is_structured: bool = False) -> MagicMock
assert second_call_args["response_format"] == SampleOutput


# The `gateway:` cases that used to live here are gone. any-llm retired the
# `gateway` provider key in favour of `otari`, so `split_model_provider` raises
# before tinyagent reaches its `if provider_name == "gateway"` strip, and `otari:`
# cannot stand in because the otari SDK is not a dependency here. The strip lives
# in the `tinyagent` package, so that behavior belongs to its test suite.
# See https://github.com/mozilla-ai/tinyagent/issues/23
@pytest.mark.parametrize(
("model_id", "expected_uses_openai"),
[
("gateway:openai:gpt-4.1-mini", True),
("gateway:anthropic:claude-3", False),
("openai:gpt-4", True),
("anthropic:claude-3", False),
],
)
def test_uses_openai_handles_gateway_provider(
model_id: str, expected_uses_openai: bool
) -> None:
def test_uses_openai(model_id: str, expected_uses_openai: bool) -> None:
config = AgentConfig(model_id=model_id)
agent: TinyAgent = AnyAgent.create(AgentFramework.TINYAGENT, config) # type: ignore[assignment]

Expand Down
7 changes: 4 additions & 3 deletions tests/unit/tools/test_unit_wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,7 +114,7 @@ def test_wrap_tool_smolagents_builtin_tools() -> None:


def test_wrap_tool_google() -> None:
from google.adk.tools import FunctionTool
from google.adk.tools.function_tool import FunctionTool

wrapper = MagicMock()
wrapper.return_value = None
Expand All @@ -125,7 +125,7 @@ def test_wrap_tool_google() -> None:


def test_wrap_tool_google_already_wrapped() -> None:
from google.adk.tools import FunctionTool
from google.adk.tools.function_tool import FunctionTool

wrapper = MagicMock()
wrapper.return_value = None
Expand All @@ -137,7 +137,8 @@ def test_wrap_tool_google_already_wrapped() -> None:


def test_wrap_tool_google_builtin_tools() -> None:
from google.adk.tools import FunctionTool, google_search
from google.adk.tools.function_tool import FunctionTool
from google.adk.tools.google_search_tool import google_search

wrapper = MagicMock()
wrapper.return_value = None
Expand Down
Loading